diff --git a/heu/library/algorithms/ou_new/FMLM.cu b/heu/library/algorithms/ou_new/FMLM.cu new file mode 100644 index 0000000..b68bc04 --- /dev/null +++ b/heu/library/algorithms/ou_new/FMLM.cu @@ -0,0 +1,19687 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define OU_T_TAU 256 // 解密指数 t 的 bit 数(p-1 的大素因子,~256 bit) +#define OU_T_EXP_LIMBS (OU_T_TAU / 64) // = 4,每个 t 占 4 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// paillier_addhomo: 密文同态加法 c̃ = FMLM(c̃₁, c̃₂) = c₁c₂R (mod n²) +// +// 数学背景(图示公式): +// AddHomo(c₁, c₂) = c₁·c₂ mod n² +// 在蒙哥马利域中等价于: +// c̃₁ = c₁·R, c̃₂ = c₂·R +// c̃ = FMLM(c̃₁, c̃₂) = c̃₁·c̃₂·R⁻¹ = c₁c₂R (即 (c₁c₂) 的蒙哥马利表示) +// +// 实现:一次 XYfixWarpVector 内核调用,全程在 GPU 上完成,无 H2D/D2H 传输。 +// +// 参数: +// batch 多项式对数(= 待处理的密文对数) +// d_c1_tilde [batch×ARR_LEN] 输入 c̃₁,设备指针(只读) +// d_c2_tilde [batch×ARR_LEN] 输入 c̃₂,设备指针(只读,不被修改) +// d_result [batch×ARR_LEN] 输出 c̃ = FMLM(c̃₁, c̃₂),设备指针 +// 允许 d_result == d_c1_tilde(原地),但不能等于 d_c2_tilde +// 其余参数 与 paillier_subhomo 保持一致(NTT 旋转因子、模数等) +// ============================================================================= +void paillier_addhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // 把 c̃₁ 加载进输出缓冲区;XYfixWarpVector 将原地把它替换为 FMLM(c̃₁, c̃₂) + if (d_result != d_c1_tilde) { + CUDA_CHECK(cudaMemcpy(d_result, d_c1_tilde, batch_bytes, + cudaMemcpyDeviceToDevice)); + } + + // 每个 block = 1 个 warp (32 线程),处理 1 个多项式;共 batch 个 block + // d_c2_tilde 在内核中仅被读取,const_cast 是安全的 + XYfixWarpVector<<>>( + d_result, const_cast(d_c2_tilde), d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +// 保留参数兼容调用处,128-bit 时不需要 p +static void ou_gen_r_prime(const uint64_t * /*ou_p_limbs17*/, int batch, + uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================================= +// ou_encrypt +// +// 功能:批量 OU 加密(正明文) +// c̃[i] = FMLM(G^{m_i}, H^{r'_i}) mod n (FMLM 域输出) +// +// 数学依据: +// OU 加密公式:c = G^m · H^r mod n +// FMLM 域语义:FMLM(Ã, B̃) = ÷B̃·R⁻¹ mod n +// 因为 Getgp/Getrn 输出 FMLM 域(乘了 R),所以: +// FMLM(G^m·R, H^r'·R) = G^m·H^r'·R mod n ✓ +// +// 三步流程: +// Step① Getgp :G 预计算表 × m_i → G^{m_i} (FMLM 域)→ d_ct_out +// Step② Getrn :H 预计算表 × r'_i → H^{r'_i}(FMLM 域)→ d_Hr(临时) +// Step③ XYfixWarpVector:d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) +// +// 调用前准备: +// generate_G_table(Modn, ou_G, h_G_table) 并上传 → d_G_table +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM +// 域,generate_G_table 输出) d_H_table [TABLE_SIZE × ARR_LEN] H +// 预计算表(FMLM 域,generate_H_table 输出) d_m_batch [batch × +// OU_EXP_LIMBS] 明文 m(base-2^64 小端序,OU_EXP_LIMBS=22) +// d_r_prime_batch [batch × OU_HR_EXP_LIMBS] 随机指数 r'(OU_HR_TAU=128 +// bit,OU_HR_EXP_LIMBS=2) d_ct_out [batch × ARR_LEN] 输出密文(FMLM +// 域) batch 明文数量 +// +// 返回:GPU 端 Step①②③ 总耗时(ms) +// ============================================================================= +float ou_encrypt(const uint64_t *d_G_table, const uint64_t *d_H_table, + const uint64_t *d_m_batch, // [batch × OU_EXP_LIMBS] + const uint64_t *d_r_prime_batch, // [batch × OU_HR_EXP_LIMBS] + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM 域) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getgp — G^{m_i} mod n(FMLM 域)→ d_ct_out + Getgp<<>>( + d_G_table, d_m_batch, OU_TAU, d_ct_out, batch, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step③:XYfixWarpVector — d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) + // 结果:c̃[i] = G^{m_i} · H^{r'_i} · R mod n(FMLM 域密文) + XYfixWarpVector<<>>( + d_ct_out, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================= +// ou_dec +// +// 功能:GPU 批量 OU 解密第一步——模幂 c^t mod n(普通域输出) +// +// OU 完整解密流程: +// Step①(本函数): c^t mod n ← FMLE_mod3_Kernel,NTT 参数为 n 模数 +// Step②(CPU 端): (result) mod p² ← 因 n=p²·q,c^t mod n 再 mod p² = c^t +// mod p² Step③(CPU 端): m = L(c') · gp_inv mod p,其中 L(x) = (x-1)/p +// +// 输入 d_c_tilde 为 FMLM 域密文 c̃ = c·R mod n; +// FMLE_mod3_Kernel 跳过 Step8(输入已是 FMLM 域), +// 保留 Step13(乘 R⁻¹,还原为普通域),输出 c^t mod n。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(FMLM 域,mod n) +// d_t_exp [batch × OU_T_EXP_LIMBS] 指数 t(压缩格式:OU_T_EXP_LIMBS=4 个 +// uint64_t, +// 小端序,位 i 在 limb[i/64] 的第 i%64 +// 位) +// d_output [batch × ARR_LEN] 输出:c^t mod n(普通域,base-2^17) +// batch 批大小 +// d_r0 [ARR_LEN] r₀ = (2^(ARR_LEN×BASE_BITS)-1) mod n +// (FMLM 单位元,即蒙哥马利域中的 1) +// 其余为 NTT 参数(n 模数,与 ou_encrypt 完全一致) +// +// 返回:GPU 端到端耗时(ms) +// ============================================================= +float ou_dec(const uint64_t *d_c_tilde, const uint64_t *d_t_exp, + uint64_t *d_output, int batch, const uint64_t *d_r0, + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * OU_T_EXP_LIMBS; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_t_exp + exp_off, OU_T_TAU, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================================ +// ou_broadcast_kernel: 将单份 src[ARR_LEN] 广播到 dst[batch × ARR_LEN] +// ============================================================================ +__global__ void ou_broadcast_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================================ +// ou_compute_p2_kernel: 单 instance CGBN 计算 p² = p * p +// 启动参数: <<<1, INV_TPI>>> +// ============================================================================ +__global__ void ou_compute_p2_kernel(cgbn_error_report_t *report, + inv_bn_mem_t *d_p2, inv_bn_mem_t *d_p) { + int instance_id = (blockIdx.x * blockDim.x + threadIdx.x) / INV_TPI; + if (instance_id != 0) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t p, p2; + cgbn_load(env, p, d_p); + cgbn_mul(env, p2, p, p); // p ≈ 1364 bit, p² ≈ 2728 bit < 4096 bit,不截断 + cgbn_store(env, d_p2, p2); +} + +// ============================================================================ +// ou_L_kernel: GPU CGBN 批量计算 L = (c^t mod p² − 1) / p +// 每 INV_TPI(=32) 线程处理一个实例 +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_L_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_L, // [batch] 输出 + inv_bn_mem_t *d_ct, // [batch] 输入:c^t(CGBN 格式) + inv_bn_mem_t *d_p2, // [1] 输入:p²(常量,所有实例共用) + inv_bn_mem_t *d_p, // [1] 输入:p (常量,所有实例共用) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t ct, p2, p_bn, tmp, L; + + cgbn_load(env, ct, d_ct + instance_id); + cgbn_load(env, p2, d_p2); // 所有实例共用同一个 p² + cgbn_load(env, p_bn, d_p); + + cgbn_rem(env, tmp, ct, p2); // tmp = c^t mod p² + cgbn_sub_ui32(env, tmp, tmp, + 1); // tmp = tmp − 1 (OU 保证 c^t ≡ 1 mod p,故 tmp ≥ 1) + cgbn_div(env, L, tmp, p_bn); // L = tmp / p (精确整除) + + cgbn_store(env, d_L + instance_id, L); +} + +// ============================================================================ +// ou_modp_kernel: GPU CGBN 批量计算 m = prod mod p +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_modp_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_m, // [batch] 输出 + inv_bn_mem_t *d_prod, // [batch] 输入:L × gp_inv mod n + inv_bn_mem_t *d_p, // [1] 输入:p(常量) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t prod, p_bn, m; + + cgbn_load(env, prod, d_prod + instance_id); + cgbn_load(env, p_bn, d_p); + cgbn_rem(env, m, prod, p_bn); + + cgbn_store(env, d_m + instance_id, m); +} + +// ============================================================================ +// ou_dec_complete: 完整 OU 解密 +// +// Step① ou_dec : d_c_tilde → d_ct_plain (c^t mod n,普通域) +// Step② format_to_cgbn_kernel : d_ct_plain → d_ct_cgbn +// Step③ ou_L_kernel : d_ct_cgbn → d_L_cgbn (L=(c^t mod +// p²−1)/p) Step④ format_from_cgbn_kernel : d_L_cgbn → d_L_b17 Step⑤ +// XYfixWarpROneVector×1 : d_gp_inv → 蒙哥马利域(原地,单份) Step⑥ +// ou_broadcast_kernel : d_gp_inv → d_gp_inv_batch(batch 份) Step⑦ +// XYfixWarpVector : d_L_b17 × d_gp_inv_batch → L×gp_inv mod n Step⑧ +// format_to_cgbn_kernel : d_L_b17 → d_prod_cgbn Step⑨ ou_modp_kernel : +// d_prod_cgbn → d_m_cgbn (mod p) Step⑩ format_from_cgbn_kernel : d_m_cgbn +// → d_m_out +// +// 返回: GPU 全流程耗时(ms,含 ou_dec 内部时间) +// ============================================================================ +float ou_dec_complete( + const uint64_t *d_c_tilde, // [batch × ARR_LEN] FMLM 域密文 + const uint64_t *d_t_exp, // [batch × OU_T_EXP_LIMBS] 指数 t + uint64_t *d_m_out, // [batch × ARR_LEN] 输出:明文 m(base-2^17) + int batch, + const uint64_t *h_p, // p(主机端,base-2^17) + const uint64_t *h_gp_inv, // gp_inv(主机端,base-2^17) + const uint64_t *d_r0, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + uint64_t *d_ctR, // CT(R²),传给 XYfixWarpROneVector + uint64_t *d_nctR // NCT(R²),传给 XYfixWarpROneVector +) { + const size_t ct_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + const size_t cgbn_bytes = (size_t)batch * sizeof(inv_bn_mem_t); + const size_t one_cgbn = sizeof(inv_bn_mem_t); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + // ── 申请设备临时内存 ───────────────────────────────────────────────────── + uint64_t *d_ct_plain = nullptr; + inv_bn_mem_t *d_ct_cgbn = nullptr; + inv_bn_mem_t *d_p_cgbn = nullptr; + inv_bn_mem_t *d_p2_cgbn = nullptr; + inv_bn_mem_t *d_L_cgbn = nullptr; + uint64_t *d_L_b17 = nullptr; + uint64_t *d_gp_inv = nullptr; + uint64_t *d_gp_inv_batch = nullptr; + inv_bn_mem_t *d_prod_cgbn = nullptr; + inv_bn_mem_t *d_m_cgbn = nullptr; + + CUDA_CHECK(cudaMalloc(&d_ct_plain, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_p_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_p2_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_L_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_b17, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_gp_inv, ARR_LEN * sizeof(uint64_t))); + CUDA_CHECK(cudaMalloc(&d_gp_inv_batch, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_prod_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_cgbn, cgbn_bytes)); + + // ── Step①: c^t mod n ──────────────────────────────────────────────────── + ou_dec(d_c_tilde, d_t_exp, d_ct_plain, batch, d_r0, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + // ou_dec 内部使用多流,需同步后才能进行后续 CGBN 操作 + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 上传 p、gp_inv;GPU 计算 p² ────────────────────────────────────────── + { + inv_bn_mem_t h_p_cgbn; + bn17_to_cgbn_mem(h_p, &h_p_cgbn); + CUDA_CHECK( + cudaMemcpy(d_p_cgbn, &h_p_cgbn, one_cgbn, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_gp_inv, h_gp_inv, ARR_LEN * sizeof(uint64_t), + cudaMemcpyHostToDevice)); + + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + ou_compute_p2_kernel<<<1, INV_TPI>>>(report, d_p2_cgbn, d_p_cgbn); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] p² 计算 CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step②: base-2^17 → CGBN ───────────────────────────────────────────── + format_to_cgbn_kernel<<>>(d_ct_cgbn, d_ct_plain, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③: L = (c^t mod p² − 1) / p ──────────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_L_kernel<<>>(report, d_L_cgbn, d_ct_cgbn, d_p2_cgbn, + d_p_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_L_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step④: CGBN → base-2^17 (L) ──────────────────────────────────────── + format_from_cgbn_kernel<<>>(d_L_b17, d_L_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑤: gp_inv 转蒙哥马利域(单份)────────────────────────────────── + // FMLM(gp_inv, R²) = gp_inv * R mod n → 蒙哥马利域,结果原地写回 d_gp_inv + XYfixWarpROneVector<<<1, 32>>>( + d_gp_inv, d_ctR, d_nctR, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑥: 广播 gp_inv_mont 到 batch 份 ──────────────────────────────── + { + const int total = batch * ARR_LEN; + const int bthreads = 256; + const int bblocks = (total + bthreads - 1) / bthreads; + ou_broadcast_kernel<<>>(d_gp_inv_batch, d_gp_inv, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step⑦: FMLM(L_plain, gp_inv_mont) = L × gp_inv mod n(普通域)────── + // inout=d_L_b17(普通域),inoutA=d_gp_inv_batch(蒙哥马利域) + // 结果 = L * gp_inv_mont * R^{-1} = L * gp_inv mod n,写回 d_L_b17 + XYfixWarpVector<<>>( + d_L_b17, d_gp_inv_batch, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑧: base-2^17 → CGBN (L × gp_inv mod n) ───────────────────────── + format_to_cgbn_kernel<<>>(d_prod_cgbn, d_L_b17, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑨: m = (L × gp_inv mod n) mod p ──────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_modp_kernel<<>>(report, d_m_cgbn, d_prod_cgbn, d_p_cgbn, + batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_modp_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step⑩: CGBN → base-2^17 (m,写入 d_m_out) ────────────────────────── + format_from_cgbn_kernel<<>>(d_m_out, d_m_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + // ── 释放临时设备内存 ───────────────────────────────────────────────────── + cudaFree(d_ct_plain); + cudaFree(d_ct_cgbn); + cudaFree(d_p_cgbn); + cudaFree(d_p2_cgbn); + cudaFree(d_L_cgbn); + cudaFree(d_L_b17); + cudaFree(d_gp_inv); + cudaFree(d_gp_inv_batch); + cudaFree(d_prod_cgbn); + cudaFree(d_m_cgbn); + + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================================ +// paillier_L_kernel: 批量计算 L(c^λ) = (c^λ − 1) / n +// 前提:c^λ ≡ 1 (mod n),故 c^λ − 1 精确被 n 整除 +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void paillier_L_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_L, // [batch] 输出:L 值(< n) + inv_bn_mem_t *d_ct, // [batch] 输入:c^λ mod n²(普通域,已在 [0,n²)) + inv_bn_mem_t *d_n, // [1] 输入:n(常量,所有实例共用) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t ct, n_bn, tmp, L; + cgbn_load(env, ct, d_ct + instance_id); + cgbn_load(env, n_bn, d_n); + + cgbn_sub_ui32(env, tmp, ct, 1); // tmp = c^λ − 1 + cgbn_div(env, L, tmp, n_bn); // L = (c^λ − 1) / n(精确整除) + + cgbn_store(env, d_L + instance_id, L); +} + +// 主机辅助函数:base-2^17(256 limbs)→ base-2^64(n64 limbs) +static void bn17_to_bn64(const uint64_t *src17, uint64_t *dst64, int n64) { + memset(dst64, 0, (size_t)n64 * sizeof(uint64_t)); + for (int i = 0; i < 256; i++) { + uint64_t val = src17[i] & 0x1FFFFULL; + int bpos = i * 17; + int widx = bpos / 64; + int boff = bpos % 64; + if (widx >= n64) break; + dst64[widx] |= (val << boff); + if (boff + 17 > 64 && widx + 1 < n64) + dst64[widx + 1] |= (val >> (64 - boff)); + } +} + +// ============================================================================ +// paillier_dec_complete: 完整 Paillier 解密 +// +// Paillier 解密公式:m = L(c^λ mod n²) · μ mod n +// L(x) = (x − 1) / n +// +// Step① paillier_dec : d_c_tilde → d_ct_plain (c^λ mod +// n²,普通域) Step② format_to_cgbn_kernel : d_ct_plain → d_ct_cgbn Step③ +// paillier_L_kernel : d_ct_cgbn → d_L_cgbn (L=(c^λ−1)/n) Step④ +// format_from_cgbn_kernel : d_L_cgbn → d_L_b17 Step⑤ +// XYfixWarpROneVector×1 : d_mu → 蒙哥马利域(单份,原地) Step⑥ +// ou_broadcast_kernel : d_mu → d_mu_batch(batch 份) Step⑦ +// XYfixWarpVector : d_L_b17 × d_mu_batch → L×μ mod n²(普通域) +// Step⑧ format_to_cgbn_kernel : d_L_b17 → d_prod_cgbn +// Step⑨ ou_modp_kernel(复用) : d_prod_cgbn → d_m_cgbn (mod n) +// Step⑩ format_from_cgbn_kernel : d_m_cgbn → d_m_out +// +// 参数说明: +// d_c_tilde [batch × ARR_LEN] 输入:FMLM 域密文 c̃ = c · R mod n² +// d_m_out [batch × ARR_LEN] 输出:明文 m(base-2^17) +// tau λ 的 bit 长度(建议 2048) +// h_lambda [256] λ(host,base-2^17) +// h_mu [256] μ = λ⁻¹ mod n(host,base-2^17) +// h_n [256] Paillier 主模数 n(host,base-2^17) +// d_negmodn/d_modn 等 NTT 参数,对应 n²(加密模数) +// d_ctR/d_nctR R² 的 CT/NCT 形式(XYfixWarpROneVector 用) +// ============================================================================ +float paillier_dec_complete( + const uint64_t *d_c_tilde, uint64_t *d_m_out, int batch, int tau, + const uint64_t *d_lambda_exp, // [batch×tau/64] DEVICE, base-2^64 + uint64_t *d_mu_dev, // [ARR_LEN] DEVICE, base-2^17 + inv_bn_mem_t *d_n_cgbn, // [1] DEVICE, CGBN + const uint64_t *d_r0, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + uint64_t *d_ctR, uint64_t *d_nctR) { + const size_t ct_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + const size_t cgbn_bytes = (size_t)batch * sizeof(inv_bn_mem_t); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + // ── 申请设备临时内存 ───────────────────────────────────────────────────── + uint64_t *d_ct_plain = nullptr; + inv_bn_mem_t *d_ct_cgbn = nullptr; + inv_bn_mem_t *d_L_cgbn = nullptr; + uint64_t *d_L_b17 = nullptr; + uint64_t *d_mu_batch = nullptr; // μ_mont broadcast(batch 份) + inv_bn_mem_t *d_prod_cgbn = nullptr; + inv_bn_mem_t *d_m_cgbn = nullptr; + + CUDA_CHECK(cudaMalloc(&d_ct_plain, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_b17, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_mu_batch, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_prod_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_cgbn, cgbn_bytes)); + + // ── Step①: c^λ mod n²(普通域)──────────────────────────────────────── + paillier_dec(d_c_tilde, d_lambda_exp, d_ct_plain, batch, tau, d_r0, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②: base-2^17 → CGBN ───────────────────────────────────────────── + format_to_cgbn_kernel<<>>(d_ct_cgbn, d_ct_plain, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③: L = (c^λ − 1) / n ──────────────────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + paillier_L_kernel<<>>(report, d_L_cgbn, d_ct_cgbn, + d_n_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[paillier_dec_complete] paillier_L_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step④: CGBN → base-2^17 (L) ──────────────────────────────────────── + format_from_cgbn_kernel<<>>(d_L_b17, d_L_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑤: μ 转蒙哥马利域(单份,原地) + // XYfixWarpROneVector: FMLM(μ, R²) = μ·R mod n² → 蒙哥马利表示 + XYfixWarpROneVector<<<1, 32>>>( + d_mu_dev, d_ctR, d_nctR, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑥: 广播 μ_mont 到 batch 份 ────────────────────────────────────── + { + const int total = batch * ARR_LEN; + const int bthreads = 256; + const int bblocks = (total + bthreads - 1) / bthreads; + ou_broadcast_kernel<<>>(d_mu_batch, d_mu_dev, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step⑦: FMLM(L_plain, μ_mont) = L·μ mod n²(普通域)──────────────── + // FMLM(A, B) = A·B·R⁻¹ mod n² + // A = L(普通域),B = μ·R(蒙哥马利域) + // 结果 = L·μ·R·R⁻¹ = L·μ mod n² + XYfixWarpVector<<>>( + d_L_b17, d_mu_batch, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑧: base-2^17 → CGBN (L·μ mod n²) ─────────────────────────────── + format_to_cgbn_kernel<<>>(d_prod_cgbn, d_L_b17, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑨: m = (L·μ mod n²) mod n ────────────────────────────────────── + // 因 L < n,μ < n,故 L·μ < n²,L·μ mod n² = L·μ,再 mod n 得最终明文 + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_modp_kernel<<>>(report, d_m_cgbn, d_prod_cgbn, d_n_cgbn, + batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[paillier_dec_complete] ou_modp_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step⑩: CGBN → base-2^17 (m,写入 d_m_out) ────────────────────────── + format_from_cgbn_kernel<<>>(d_m_out, d_m_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + // ── 释放临时设备内存 ───────────────────────────────────────────────────── + cudaFree(d_ct_plain); + cudaFree(d_ct_cgbn); + cudaFree(d_L_cgbn); + cudaFree(d_L_b17); + cudaFree(d_mu_batch); + cudaFree(d_prod_cgbn); + cudaFree(d_m_cgbn); + + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +int main() { + srand(time(NULL)); + const int LIMBS = 256; + const int POLY_PER_STREAM = 16384; + const int m = 4; + const int NUM_WARMUP = 5; // 预热轮次(不计时) + const int NUM_TIMED_RUNS = 30; // 正式计时轮次 + + int total_polynomials = NUM_STREAMS * POLY_PER_STREAM; + size_t total_elements = (size_t)total_polynomials * LIMBS; + size_t total_bytes = total_elements * sizeof(uint64_t); + + // ======================================================== + // 1. Host 内存分配 + 随机输入生成 + // ======================================================== + uint64_t *arrX, *arrY; + cudaMallocHost(&arrX, total_bytes); + cudaMallocHost(&arrY, total_bytes); + + uint64_t mask_17bit = (1ULL << 17) - 1; + uint64_t mask_16bit = (1ULL << 16) - 1; + + for (int s = 0; s < NUM_STREAMS; s++) { + for (int p = 0; p < POLY_PER_STREAM; p++) { + int offset = (s * POLY_PER_STREAM + p) * LIMBS; + for (int j = 0; j < LIMBS; j++) { + if (j < 240) { + arrX[offset + j] = rand() & mask_17bit; + arrY[offset + j] = rand() & mask_17bit; + } else if (j == 240) { + arrX[offset + j] = rand() & mask_16bit; + arrY[offset + j] = rand() & mask_16bit; + } else { + arrX[offset + j] = 0; + arrY[offset + j] = 0; + } + } + } + } + + // ======================================================== + // 2. 备份原始 arrX + // ======================================================== + uint64_t *arrX_backup; + cudaMallocHost(&arrX_backup, total_bytes); + memcpy(arrX_backup, arrX, total_bytes); + + // ======================================================== + // 3. 常量内存池(Host → Device,计时区外只做一次) + // ======================================================== + const int NUM_ARRAYS = 13; + size_t pool_bytes = NUM_ARRAYS * LIMBS * sizeof(uint64_t); + + uint64_t *h_pool; + cudaMallocHost(&h_pool, pool_bytes); + + memcpy(h_pool + 0 * LIMBS, con_negmodn, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 1 * LIMBS, con_NegModn_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 2 * LIMBS, con_modn, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 3 * LIMBS, con_modn_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 4 * LIMBS, con_twiddle, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 5 * LIMBS, con_twiddle_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 6 * LIMBS, con_ICTTwiddle, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 7 * LIMBS, con_ICTTwiddle_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 8 * LIMBS, con_twiddle_NCT, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 9 * LIMBS, con_twiddle_NCT_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 10 * LIMBS, con_InvTwiddle, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 11 * LIMBS, con_InvTwiddle_shoup, LIMBS * sizeof(uint64_t)); + memcpy(h_pool + 12 * LIMBS, con_sample, LIMBS * sizeof(uint64_t)); + + uint64_t *d_pool; + cudaMalloc(&d_pool, pool_bytes); + cudaMemcpy(d_pool, h_pool, pool_bytes, cudaMemcpyHostToDevice); + + CommonDeviceArrays common; + common.n = 256; + common.mod = 25668312996353ULL; + common.inv = 25568046148711ULL; + common.inv_shoup = 18374686479671626487ULL; + + common.d_negmodn = d_pool + 0 * LIMBS; + common.d_negmodn_shoup = d_pool + 1 * LIMBS; + common.d_modn = d_pool + 2 * LIMBS; + common.d_modn_shoup = d_pool + 3 * LIMBS; + common.d_twiddle = d_pool + 4 * LIMBS; + common.d_twiddle_shoup = d_pool + 5 * LIMBS; + common.d_ICTTwiddle = d_pool + 6 * LIMBS; + common.d_ICTTwiddle_shoup = d_pool + 7 * LIMBS; + common.d_NCTtwiddle = d_pool + 8 * LIMBS; + common.d_NCTtwiddle_shoup = d_pool + 9 * LIMBS; + common.d_InvNCTtwiddle = d_pool + 10 * LIMBS; + common.d_InvNCTtwiddle_shoup = d_pool + 11 * LIMBS; + common.d_sample = d_pool + 12 * LIMBS; + + // ======================================================== + // 4. 设备数据内存(计时区外一次性分配) + // ======================================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 初始上传(计时区外) + cudaMemcpy(d_arrX, arrX_backup, total_bytes, cudaMemcpyHostToDevice); + cudaMemcpy(d_arrY, arrY, total_bytes, cudaMemcpyHostToDevice); + + // ======================================================== + // 5. 创建流 + 计算 grid/block 参数 + // ======================================================== + int grid_size = (POLY_PER_STREAM + m - 1) / m; + int threads_per_block = m * 32; + + cudaStream_t *streams = new cudaStream_t[NUM_STREAMS]; + for (int i = 0; i < NUM_STREAMS; i++) cudaStreamCreate(&streams[i]); + + printf("[配置] %d 流 x %d 多项式 = %d 模乘/轮,grid=%d blocks=%d threads\n", + NUM_STREAMS, POLY_PER_STREAM, total_polynomials, grid_size, + threads_per_block); + + // 用于简化内核调用的宏(避免参数列表重复书写) +#define LAUNCH_ALL_STREAMS() \ + do { \ + for (int _i = 0; _i < NUM_STREAMS; _i++) { \ + int _off = _i * POLY_PER_STREAM * LIMBS; \ + XYfixWarpVector<<>>( \ + &d_arrX[_off], &d_arrY[_off], common.d_negmodn, \ + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, \ + common.n, common.mod, common.d_twiddle, common.d_twiddle_shoup, \ + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, \ + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, \ + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, \ + common.d_sample); \ + } \ + } while (0) + + // ======================================================== + // 6. 预热(不计时):消除 GPU 频率爬坡、cache 冷启动 + // ======================================================== + printf("[预热] 开始 %d 轮预热...\n", NUM_WARMUP); + for (int w = 0; w < NUM_WARMUP; w++) { + LAUNCH_ALL_STREAMS(); + } + cudaDeviceSynchronize(); + printf("[预热] 完成\n\n"); + + // 预热后恢复初始输入,保证计时从已知状态出发 + cudaMemcpy(d_arrX, arrX_backup, total_bytes, cudaMemcpyHostToDevice); + cudaDeviceSynchronize(); + + // ======================================================== + // 7. 精确计时 + // + // 原理:利用 Legacy Default Stream(stream 0)的隐式屏障特性: + // ① ev_start 记录在 stream 0: + // 所有后续非默认流操作等待 ev_start 完成才能开始 + // ② ev_stop 记录在 stream 0: + // stream 0 等待所有非默认流操作全部完成后才记录时间戳 + // 因此 ev_stop - ev_start 精确捕获了所有并发 kernel 的总执行时间, + // 不含 H2D / D2H 传输开销,也不含 kernel launch 的 CPU 提交延迟。 + // ======================================================== + cudaEvent_t ev_start, ev_stop; + cudaEventCreate(&ev_start); + cudaEventCreate(&ev_stop); + + cudaEventRecord(ev_start, 0); // ← 屏障:非默认流 kernel 等待此事件 + for (int run = 0; run < NUM_TIMED_RUNS; run++) { + LAUNCH_ALL_STREAMS(); + } + cudaEventRecord(ev_stop, 0); // ← stream 0 等待所有流 kernel 完成 + cudaDeviceSynchronize(); // CPU 等待 ev_stop 记录完毕 + + float total_ms = 0.0f; + cudaEventElapsedTime(&total_ms, ev_start, ev_stop); + + long long total_muls = (long long)NUM_TIMED_RUNS * total_polynomials; + double us_per_mul = (double)total_ms * 1000.0 / (double)total_muls; + double throughput = (double)total_muls / ((double)total_ms * 1e-3); // 次/秒 + + printf("========================================\n"); + printf(" FMLM 模乘基准测试结果\n"); + printf("========================================\n"); + printf(" 预热轮次 : %d\n", NUM_WARMUP); + printf(" 计时轮次 : %d\n", NUM_TIMED_RUNS); + printf(" 每轮模乘数 : %d\n", total_polynomials); + printf(" 总模乘数 : %lld\n", total_muls); + printf(" Kernel 总执行时间 : %.3f ms\n", total_ms); + printf(" 平均单次模乘耗时 : %.4f us\n", us_per_mul); + printf(" 吞吐量 : %.2f M次/s\n", throughput / 1e6); + printf("========================================\n\n"); + +#undef LAUNCH_ALL_STREAMS + + // ======================================================== + // 8. 释放资源 + // ======================================================== + cudaEventDestroy(ev_start); + cudaEventDestroy(ev_stop); + + for (int i = 0; i < NUM_STREAMS; i++) cudaStreamDestroy(streams[i]); + delete[] streams; + + cudaFree(d_arrX); + cudaFree(d_arrY); + cudaFreeHost(arrX_backup); + cudaFreeHost(arrX); + cudaFreeHost(arrY); + cudaFreeHost(h_pool); + cudaFree(d_pool); + cudaDeviceReset(); + + return 0; +} diff --git a/heu/library/algorithms/ou_new/ou_addandsubp.cu b/heu/library/algorithms/ou_new/ou_addandsubp.cu new file mode 100644 index 0000000..3915081 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_addandsubp.cu @@ -0,0 +1,19445 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 11687101653036, 18446743758138104250, 14211471509222940197, + 4235272491137479070, 18240298465959804861, 3612012136679214338, + 17250142769850486366, 16237778209286303069, 4059562170797630074, + 9869391574692501418, 11541833079389403685, 6455235399891900791, + 12385636973677184651, 16642176613718519917, 18371701685023660666, + 12908184103137142305, 432025558415355680, 13421505415041197067, + 3981525027614746084, 11947396656646251748, 17803480186391939853, + 8964175857302805981, 10755671907831908778, 1904820253353466484, + 10527642910445979147, 13113465356235953778, 7234975774234226545, + 9757071334591424272, 7316737729072523157, 6883729646458586186, + 8932421558757596044, 14597307983516136610, 14233691893978885667, + 11827243429966767869, 4500161214914298038, 10200789522299270510, + 13311736198424319672, 15003914107128413955, 15970198162388647651, + 5178000660069144055, 12257227048244175003, 199311466455739912, + 1199638074941611369, 7733792994982443122, 5589666049742506788, + 13804186403958915190, 17844357954872068228, 1608648031291287388, + 14833453796159433687, 17508515457688533059, 5757945642895465137, + 13081582882382389324, 5394527028006066918, 276092650195297071, + 17347510335080686092, 1701269284563161833, 13303804442297711418, + 8253121806998843455, 1803714610749342533, 15051344875346831329, + 17235261944528818002, 8632334691816709691, 7259437303239191782, + 15692170915486480673, 16909097158836466193, 3813100579756643340, + 8120672335331207645, 17658082942068012540, 14712625527555672008, + 3013490167685507391, 17053781224112072993, 7951833678156564288, + 16882129623470333590, 16598623833219974769, 8844055318436626140, + 7109452390029814739, 17973202994978822509, 17988667507327805934, + 4123463954687411915, 4940157814071223716, 10034667236662981698, + 2823422309362454166, 8801094461519448259, 16352417325266577247, + 4120137837347666316, 13914619055383639681, 13924524887541743843, + 16792129803426718801, 17665685411675475587, 4424050859418317508, + 7334727500762095896, 17503735803347907081, 17481354675564559188, + 3960789173235012097, 2382836270529746228, 4344431928835757414, + 10028103253977297380, 10742172769839739841, 4046021704283408275, + 15048763124930553406, 12884462872804057858, 6406919215710243952, + 2305936473049040111, 13121402678735537817, 4908133312471037493, + 11377681924055600462, 12660562847764485932, 4317962794267615766, + 1681209879049333151, 17555975506207345043, 5125753427319466102, + 3335447880892075219, 8356915887911857374, 8584812879987417990, + 17895049134452250854, 9085212301170466888, 7267866673654176139, + 11010976396069557933, 3608178248288276404, 676841772753514930, + 14867830803115014463, 1834280874555657868, 16310358636794094835, + 7330665673582977596, 15741791681143414831, 210676699798228406, + 17198551982049102727, 6109711417879930926, 10546103406410231004, + 15078747006884146927, 15249364729241398859, 7659463200845052268, + 6442795927427660859, 2250605931405808055, 8092318475578226474, + 18259756362431830946, 17518863421902388405, 10430337473484617894, + 11467857499548239759, 14850024957617392139, 4520997378651548347, + 1002071158803202119, 13705130616563662128, 16248739479905188442, + 3181222670542328118, 10465632553806001024, 13994891411389346011, + 10984874395183979750, 9226503659276477277, 16804196055515371214, + 10159135197864231902, 2843438330111370559, 10398977526462394360, + 13629959477406403811, 17564676539751491213, 9169917922917002660, + 9102739415085474595, 2571556159645270064, 15524480249380336561, + 16515752187799014418, 15631770625703612855, 15987278742587054296, + 6287084574542908707, 5110496278831165043, 3153708236541173992, + 12927769102029794814, 15247894309294568262, 9307521059673973398, + 2770646931232435574, 16454025378125412919, 10633343977751814089, + 1784373777963240394, 6934373261888962825, 17349218505508302052, + 10968015619489026254, 1752595997967202555, 12837659792768865327, + 14742784040235182007, 10962127976630829136, 5900158086628228693, + 17316277940499271723, 348494413055888541, 488358709608098659, + 10382164829144562111, 10796492275262689178, 10517686393808888673, + 779545184377220173, 3381663347793212032, 15236282164919686367, + 6334307549334192514, 3063003522052060686, 12114810039050492787, + 8870556826759400006, 4038453701720268007, 14314379071943608840, + 4339980657355467038, 9171890896160995321, 7917821014549284449, + 13180571956635383946, 18248798750186441796, 14577763713235528379, + 1809799949882400184, 6318379507872769143, 1709904138639204836, + 4433595655137503557, 14198791200961044293, 4650959752702308354, + 3318171583106947229, 209694021303727738, 1076839995989163599, + 16905707527716642672, 3695319746732430913, 10252674333030518371, + 16420700828256874086, 2433845634309674122, 16595131843996099370, + 16829576163922136336, 4841410023332049560, 12592434475652961886, + 17572224200405592702, 2431385938003789647, 10061028979483934059, + 2925075822122586101, 8606434114160323337, 5607119374066730577, + 11884780541782053893, 3126131661631420656, 10027052590555524485, + 16048091853305621015, 15852396680435263215, 17385109108245871297, + 12899005442699559936, 7848549015331456901, 9096729734807002481, + 4996467004929486051, 12243245730936161727, 9057745574396783344, + 7771655314147204603, 17881871823990369590, 15212325419875966733, + 11042754214829301385, 3281380445998958437, 17088850123667971831, + 14495632125498788619, 17994272834268936450, 16829150404837372, + 6402480299793703941, 10533393325012975763, 2416806924432625423, + 2875142845742952402, 12269921477737603466, 712826757316093295, + 4415075740707273176, 7975119161839365838, 11673666813117204770, + 9840315584201792241}; + +const uint64_t con_modn_shoup[256] = { + 14252880640204352951, 18322338132012530560, 6849211572557864124, + 15932334426407894751, 14177538223384229049, 9106170071204927292, + 16758578487033067404, 5864117390783715525, 16494545899140440021, + 5897258617717822505, 1933174352700629569, 12810083258791009448, + 12690514865985841899, 3970720354745169798, 4239814533413767079, + 5609102486863112468, 11230284723594595426, 17034417004294615591, + 5986132948557996967, 1868566874544028188, 2158239585541928173, + 17841097290863850509, 9305374060203424222, 10083694270531949160, + 2654649954734757684, 17721823101598855865, 599980504891318557, + 17732121566266018424, 1832524248816892725, 5295674783104160211, + 13283213776815003617, 10900691717351424196, 6021057974446928650, + 12624795618036119368, 2798162278377969124, 7399862538297940480, + 576839721297220460, 16060704215571397670, 16380205270440154947, + 7499979448419237887, 13254841413012858481, 10664669973443882596, + 1765312882701300737, 9426266066319221551, 12823009007753704482, + 10630699336822577567, 16298120910453621338, 13674950148586695572, + 17678273253225120972, 6798806775207395115, 13410427498759750653, + 2614783784077562964, 9342414901595102647, 14373786281595714575, + 6330183866169305354, 998675938268783033, 10221732541776071598, + 8979762078911881033, 9878667596621300249, 4856285279936479033, + 14833025980849776526, 5604878655902262946, 13088552648421720803, + 1801154013367199521, 4158119999782205988, 7904891504652660262, + 9042945763108429841, 4642264688478488779, 16204979912313018920, + 3580705517878336362, 9712754433060271621, 8675179099278674516, + 8186655178929728093, 1884659203003867161, 17775229374523385263, + 6390348527000753038, 6439058351892770174, 2339745637453507323, + 16274314407512660647, 2247518004490005028, 18003796786185432156, + 4540940947376355923, 11987538437574474975, 16166798012420960901, + 16121611900792272328, 16082928115738878740, 121528093926685229, + 11609994994605905995, 2593441955413327993, 16920803883743198476, + 11945409668615507125, 15459882499135165139, 4709903422099897132, + 1915945056478813527, 17487099108173624447, 4121351438621439846, + 11648490996845515622, 16906896413860707859, 440932069689474224, + 5596373384320545758, 6286719224840257488, 15070666469307485122, + 1718056780076659255, 16292491121877970301, 16399246121914763003, + 868264559834958645, 5880650461523548368, 13037697811177873232, + 12598280349103069353, 8787439026840426841, 75102682845531848, + 10793543124523682506, 4058772666671965704, 4575391113880810276, + 7977675084418792789, 3637051392050280908, 16362683407568863478, + 18347383388798481077, 9115743514592553391, 9421569851894468249, + 15594101773322529942, 11807267208082355523, 2211845703086696074, + 17348335706114235958, 11847926503719721254, 17547278040398801999, + 5869056178242350580, 13320003773654588467, 12699478824066277151, + 5239882100070320334, 7261256595809982529, 18328110662323898796, + 14262563528151153433, 10570694000294503318, 12813828885908106934, + 10929763919809758798, 2938308820079533814, 12010181661893483546, + 5724066348617601097, 14693589767406584560, 7346156359909105313, + 12463844683080763521, 6157213913132141689, 10056871538135474507, + 10533920527198962255, 8235152268042627670, 12319030087627737531, + 8540756947882872704, 7431325835456426550, 1301653294393655025, + 5378735902526386334, 14912383613060771114, 639721130790109100, + 16570337161183830109, 3674985562098081097, 25882515425358888, + 5781372063417524209, 16445334884700810773, 3544553957273777187, + 7670642182104980993, 9626872654279745485, 16105190074590295359, + 7770490841006776992, 15228876210409060859, 9849662374993193055, + 8654391918266496929, 6489400423788825431, 3633531112627186925, + 14858949636521671041, 1105232854426717343, 16217154593325743207, + 6106199821319202195, 15821125396653981131, 397434568115144159, + 5468761408652955936, 1296217405573136392, 8004677824854586354, + 2227606275875858042, 810603102045699119, 10604814613007946604, + 5290458938805336352, 17851909192937068431, 13268718299834334195, + 10806219279687765004, 7326952643401977865, 5984244256621982617, + 1659224770285885321, 10142490388121661620, 543184966312697096, + 6161334213393132777, 8606758596526178508, 16789552215120228336, + 8023727110355433065, 4647377966945102884, 6753109868947696463, + 3601294586865406454, 17607120286020078979, 2828322973879047780, + 2912791741152772408, 12563760677479960004, 14534280132691080513, + 1458078075305867271, 965960906590360464, 7718560025107401880, + 5982496231126478023, 9871084626240060187, 13176440103626612310, + 12254705932735020033, 6091017959275360856, 8575903195682037371, + 12248661790925770989, 15874428453561837902, 11580211822667122432, + 1675684581791044528, 953808119473653258, 5212502010992923248, + 11653707338854681558, 2814312282799796103, 10741977896772006037, + 792456711582471131, 5829394638712031862, 14582582386261703791, + 15116195068952376181, 16192690152961596921, 14996186982344279167, + 12452747948715197047, 9822110408961686525, 13084672463213903010, + 15511412028972873361, 4378034570765898628, 7434337709561426193, + 10855303285220731112, 7759227917166922418, 4976939851858003292, + 5204453497818198107, 16768838491807377833, 9561958848337334713, + 703798444570847189, 14816217796224051499, 2968875028840941609, + 6519175664834707328, 17450997194073458375, 12811758118865484166, + 10759621678827047865, 11859099579116844022, 14425180740705110568, + 7511257586748720299, 3539736587530686592, 4447312216206097890, + 11184913710261542236, 8771977661991653642, 12354338902272316516, + 15541265547581959679, 14587017360825887711, 15248329331280087903, + 10992628732632760992}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 15 +#define TABLE_SIZE (1 << WINDOW_BITS) // 2048 个表项,对应 11-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + // ── 优化①:扫描最高非零窗口,跳过前导零的冗余迭代 ────────────────────────── + // window_val 由 __shfl_sync 推导,对 warp 内所有 lane 完全一致, + // 因此下面的 if/break 不产生 warp 分支分歧。 + int k_start = -1; + uint64_t first_wval = 0; + for (int k = num_windows - 1; k >= 0; k--) { + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t wv = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + wv |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + if (wv != 0) { + k_start = k; + first_wval = wv; + break; + } + } + + // p == 0:直接输出 FMLM 恒等元 table[0],无需任何模幂运算 + if (k_start < 0) { +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = __ldg(d_G_table + 8 * lane_id + j); + return; + } + + // 用最高非零窗口直接初始化累加器,跳过对恒等元的冗余平方 + const uint64_t *init_entry = d_G_table + first_wval * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = __ldg(init_entry + 8 * lane_id + j); + __syncwarp(); + + // ── 主循环:从 k_start-1 向下,只处理有效窗口 ──────────────────────────── + for (int k = k_start - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ── 优化②:window_val == 0 时乘以恒等元结果不变,直接跳过 ────────── + // window_val 对整个 warp 一致,if 不产生分支分歧 + if (window_val != 0) { + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = __ldg(entry + 8 * lane_id + j); + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + // ── 优化①:扫描最高非零窗口,跳过前导零的冗余迭代 ────────────────────────── + int k_start = -1; + uint64_t first_wval = 0; + for (int k = num_windows - 1; k >= 0; k--) { + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t wv = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + wv |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + if (wv != 0) { + k_start = k; + first_wval = wv; + break; + } + } + + // p == 0:直接输出 FMLM 恒等元 table[0],无需任何模幂运算 + if (k_start < 0) { +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = __ldg(d_invG_table + 8 * lane_id + j); + return; + } + + // 用最高非零窗口直接初始化累加器,跳过对恒等元的冗余平方 + const uint64_t *init_entry = d_invG_table + first_wval * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = __ldg(init_entry + 8 * lane_id + j); + __syncwarp(); + + // ── 主循环:从 k_start-1 向下,只处理有效窗口 ──────────────────────────── + for (int k = k_start - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ── 优化②:window_val == 0 时乘以恒等元结果不变,直接跳过 ────────── + if (window_val != 0) { + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = __ldg(entry + 8 * lane_id + j); + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +int main() { + uint64_t Modn[256] = { + 43351, 84159, 17007, 126963, 115814, 64975, 14865, 122878, 58093, + 76773, 16638, 87086, 110462, 105466, 35053, 36095, 8051, 116177, + 119699, 118157, 12357, 71314, 68424, 35266, 58013, 63468, 22117, + 10903, 124058, 90359, 68490, 117774, 56449, 45990, 26837, 86153, + 120741, 31603, 78596, 24019, 45134, 33649, 61458, 59406, 88868, + 60745, 113313, 123484, 30017, 98185, 93108, 73040, 39521, 18181, + 2647, 51647, 10194, 73702, 22934, 64, 29664, 94536, 9414, + 63827, 6028, 107137, 71399, 49216, 8196, 46100, 117329, 67195, + 25041, 122567, 110161, 82524, 85064, 85420, 38367, 90728, 6216, + 87366, 124652, 29067, 100922, 38894, 64688, 22860, 83774, 130371, + 39036, 94816, 45277, 76221, 67984, 78245, 70889, 64430, 52640, + 50933, 54580, 32496, 95587, 110988, 102834, 68631, 42744, 111149, + 127114, 116295, 108662, 4710, 31837, 15424, 50234, 99229, 61393, + 81585, 33195, 14128, 9168, 55047, 119038, 97329, 43164, 111637, + 39396, 13009, 90209, 92184, 81272, 101938, 57149, 82121, 100630, + 37780, 7881, 13181, 8505, 125111, 43862, 119168, 19431, 80034, + 114187, 71294, 52911, 81495, 14533, 87246, 126978, 30310, 9978, + 44551, 60081, 126942, 75376, 77030, 36034, 104993, 58885, 90371, + 111023, 45378, 97203, 126393, 72942, 8192, 124336, 37338, 116797, + 66693, 60337, 12040, 90738, 108119, 66171, 78981, 79494, 91989, + 89494, 118041, 29798, 30883, 110522, 122729, 7823, 62523, 20666, + 52089, 43045, 51146, 24317, 38753, 122735, 100047, 56716, 117911, + 60032, 27220, 44093, 56, 113553, 49629, 84418, 64845, 67097, + 10050, 8296, 90055, 75973, 63190, 82919, 56713, 30800, 90227, + 63208, 39501, 61899, 129744, 78395, 58460, 121961, 72489, 20054, + 64673, 102069, 663, 109348, 36701, 18676, 98666, 77634, 108307, + 103092, 49932, 49372, 39729, 72337, 90146, 247, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49877, 103520, 14633, 105196, 88121, 75472, 97344, 77893, 104350, + 114542, 91892, 50952, 64070, 21857, 78639, 2966, 6840, 35973, + 24928, 25085, 49586, 41185, 13904, 75536, 2710, 53110, 104559, + 109441, 40809, 28751, 119672, 47293, 31768, 78934, 55525, 52600, + 69445, 27490, 96611, 84604, 11133, 106316, 35838, 34416, 127179, + 34651, 118161, 16314, 45331, 107582, 24235, 12592, 81941, 19278, + 85161, 104589, 13239, 111883, 81763, 64318, 45287, 80867, 37986, + 126633, 118333, 32371, 49588, 111307, 6225, 98536, 50417, 128492, + 7707, 81120, 51703, 13989, 14502, 86184, 74055, 95503, 80070, + 61934, 73173, 322, 128915, 100622, 92460, 3170, 90102, 44305, + 79192, 110628, 84896, 92155, 13698, 129632, 82311, 123141, 99527, + 28216, 52443, 78019, 37062, 11803, 15622, 2177, 42554, 81945, + 37634, 97471, 11261, 30170, 62893, 129991, 77778, 123677, 75667, + 22518, 67693, 79122, 34634, 117067, 33393, 41664, 46880, 63508, + 1202, 120484, 35970, 14884, 64323, 87199, 61041, 17700, 6496, + 128561, 104072, 129267, 31935, 119451, 54205, 55539, 59285, 49953, + 34036, 45769, 28494, 30664, 79616, 47756, 57833, 99317, 122074, + 130764, 65477, 42562, 43963, 97923, 114657, 34735, 19193, 84201, + 24821, 98149, 101368, 100860, 110862, 96240, 101375, 55675, 99994, + 12323, 56026, 120364, 84030, 10327, 108568, 4795, 122128, 3775, + 1740, 113334, 5740, 61052, 15255, 44939, 84950, 4631, 87197, + 63464, 45041, 49844, 102052, 41710, 76594, 96748, 2213, 15033, + 56862, 42121, 22702, 29104, 55955, 32193, 62378, 61812, 37549, + 27929, 118796, 116386, 35884, 83278, 116744, 103768, 106752, 29801, + 6976, 81713, 55669, 12038, 51733, 6915, 128541, 82038, 102167, + 64630, 125581, 69829, 79662, 80895, 89416, 41571, 113918, 73736, + 22655, 72892, 97009, 75512, 83469, 50798, 35893, 72631, 27752, + 114176, 116066, 35199, 11556, 117400, 53979, 71662, 76589, 35790, + 51797, 38295, 48839, 44050}; + uint64_t R1[256] = { + 95787, 87855, 74590, 120089, 96462, 125333, 89873, 62820, 56744, + 93675, 114260, 58407, 55044, 4742, 20922, 129032, 18634, 103071, + 2852, 114517, 116272, 79216, 95365, 35177, 128432, 80425, 18923, + 592, 13588, 42144, 48019, 39668, 66805, 42663, 33194, 65911, + 93428, 9610, 76041, 48300, 121686, 67062, 30099, 86626, 99273, + 90908, 66468, 28073, 71719, 43868, 40340, 88274, 109318, 21824, + 16472, 116161, 1346, 106033, 20342, 35258, 20632, 105594, 118266, + 97653, 97643, 7306, 67863, 79950, 112151, 117205, 39906, 100559, + 19328, 14826, 43881, 64539, 123341, 37113, 15909, 43631, 27755, + 12868, 87791, 110907, 2763, 41576, 76238, 21079, 28709, 67173, + 22692, 45867, 80137, 111080, 90017, 19215, 15056, 7843, 34411, + 10397, 47157, 113197, 77959, 43337, 123310, 26898, 103324, 95568, + 37773, 53742, 58532, 64900, 12429, 109482, 75505, 70429, 89935, + 67404, 103144, 45403, 28839, 100826, 80183, 60279, 60825, 67114, + 15456, 95163, 5820, 106812, 38605, 127798, 43023, 23037, 109334, + 82354, 36764, 29882, 1460, 109709, 70002, 61938, 129339, 12574, + 23578, 116415, 67219, 51854, 56951, 3482, 98043, 101818, 116934, + 91679, 101444, 12152, 24202, 116763, 23200, 65030, 72892, 11401, + 89306, 104758, 94473, 69024, 2331, 120575, 38998, 71116, 123520, + 12246, 31511, 15417, 54824, 35449, 51579, 129451, 119392, 117118, + 124644, 45532, 83498, 101978, 4264, 7165, 33903, 43033, 21694, + 89206, 4834, 104846, 13789, 111463, 68368, 35303, 42736, 93343, + 26110, 67668, 100827, 18175, 80987, 95026, 43049, 29875, 73340, + 103134, 127164, 12263, 68500, 44986, 40419, 118430, 83125, 27361, + 83899, 599, 102652, 96223, 98631, 95859, 60979, 40494, 118153, + 71598, 113945, 125768, 2378, 115813, 13443, 7038, 3598, 53377, + 26625, 110573, 3294, 119847, 72218, 85832, 125, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 2023, 29313, 55760, 30415, 117209, 68859, 126088, 105595, 44062, + 129614, 108803, 62812, 119884, 130526, 4554, 108683, 53422, 114202, + 34719, 66446, 65067, 23670, 9591, 82680, 20040, 77609, 34903, + 85401, 37760, 43899, 14395, 67595, 12710, 101806, 77200, 103316, + 65099, 125953, 105217, 67808, 5979, 63677, 12355, 111674, 70131, + 45594, 117232, 10425, 81356, 112691, 14758, 10490, 24556, 27922, + 32980, 20928, 118420, 17204, 4244, 126937, 116849, 106497, 51321, + 114935, 45503, 15461, 59271, 111583, 30113, 103352, 10622, 32510, + 41116, 86928, 10137, 101567, 30707, 124356, 108755, 8203, 11158, + 9603, 114740, 5093, 13054, 61800, 75687, 38080, 11550, 87153, + 33247, 18929, 66437, 13511, 39575, 18765, 61155, 77315, 112366, + 76906, 125693, 40793, 40582, 43161, 81338, 111531, 84813, 49322, + 71309, 83250, 14948, 44745, 13967, 98243, 116072, 5842, 82567, + 77993, 80649, 107659, 66320, 122438, 54394, 54983, 79006, 105681, + 94840, 79085, 41950, 106863, 130420, 89427, 83726, 86511, 44750, + 12837, 47751, 33678, 115313, 66053, 43941, 100068, 107956, 169, + 62673, 70167, 106884, 28460, 37125, 129538, 93387, 56010, 35558, + 40621, 77463, 39765, 16013, 85203, 39465, 19946, 122865, 58068, + 60861, 54036, 48234, 130529, 59321, 83170, 116672, 11733, 94357, + 35207, 46800, 22992, 47306, 80973, 28136, 59828, 4338, 109019, + 16604, 58136, 123247, 75151, 11948, 105570, 86992, 72951, 15873, + 12088, 27387, 60538, 80333, 72819, 97245, 41421, 126078, 82127, + 78747, 70199, 99633, 38926, 56088, 20821, 130509, 125699, 127559, + 47166, 102354, 71946, 79468, 81253, 11120, 37016, 20384, 95174, + 126298, 33756, 60551, 106592, 43758, 59029, 22997, 63753, 65367, + 47855, 5100, 118402, 102326, 108278, 49481, 99388, 114710, 101475, + 124482, 27406, 78309, 52397, 69291, 787, 10, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ========================================================================= + // ou_addAndsubp 正确性测试 + // + // 密文加明文:c̃_out = FMLM(c̃, G^p·r mod n) = c·G^p·r mod n(FMLM 域) + // 密文减明文:c̃_out = FMLM(c̃, invG^p·r mod n) = c·invG^p·r mod n(FMLM 域) + // + // 流程: + // 1. 定义公钥 ou_G / ou_invG,CPU 端建立预计算表,上传 GPU + // 2. 随机生成 BATCH 个密文(< n,接近 4096 bit) + // 和明文(< p,接近 1363 bit) + // 3. H2D 传输密文×2(ADD/SUB 各一份)+ 明文,计时 + // 4. GPU:ou_addAndsubp(is_add=true) → d_ct_add 原地覆写 + // 5. GPU:ou_addAndsubp(is_add=false) → d_ct_sub 原地覆写 + // 6. D2H 传输两份结果,计时 + // 7. 写入 ou_addAndsubp_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 1000000; + + // OU 公钥 G(base-2^17,小端序,256 limb) + static const uint64_t ou_G[ARR_LEN] = { + 106476, 43971, 128828, 93344, 6696, 19006, 123835, 84284, 130736, + 17990, 26048, 72490, 70614, 22448, 108454, 62553, 78468, 81393, + 122603, 77178, 45293, 60274, 55433, 87803, 76771, 81109, 98247, + 122300, 22754, 3679, 17780, 109494, 129489, 73701, 33577, 85290, + 61919, 106717, 33528, 28665, 94695, 45005, 68911, 70333, 88542, + 58868, 86785, 83397, 11510, 116629, 28445, 96533, 126070, 72616, + 67335, 11492, 122574, 40178, 72210, 63322, 121381, 24191, 69914, + 32116, 44664, 20944, 68975, 25088, 58967, 116052, 60547, 120960, + 119878, 70784, 90289, 100265, 45899, 4808, 10773, 96336, 70004, + 40632, 74961, 82572, 102583, 92231, 57069, 79015, 107028, 66308, + 80179, 74690, 23866, 25975, 71759, 115443, 77156, 55444, 33714, + 85458, 17771, 36056, 101050, 126552, 107704, 5243, 18837, 9892, + 9282, 129831, 111243, 6831, 3584, 40846, 119112, 15226, 46947, + 41405, 13045, 22014, 127339, 74943, 105622, 76032, 89860, 60229, + 118237, 64311, 129625, 43238, 79135, 8697, 107206, 27638, 43638, + 102812, 126121, 53006, 17406, 67824, 29297, 120890, 104418, 94593, + 19381, 23189, 12753, 129574, 35126, 130937, 108605, 80367, 77402, + 112351, 99669, 74942, 8710, 94873, 8797, 112311, 45381, 125564, + 109720, 123588, 49172, 104108, 3259, 31317, 97577, 7968, 130968, + 88371, 43213, 45584, 32536, 17547, 47469, 32610, 49049, 9, + 43949, 71808, 61643, 90674, 46892, 87542, 70431, 99776, 99858, + 23393, 57574, 127452, 59500, 102555, 47948, 54920, 41570, 5626, + 91020, 76667, 62152, 8807, 129112, 70381, 88792, 35502, 47837, + 35589, 5298, 74076, 74434, 38127, 86697, 16215, 54784, 5803, + 101168, 130729, 121263, 16663, 14759, 92283, 32919, 17980, 39580, + 60077, 80927, 19921, 42631, 87624, 6138, 90365, 4191, 76995, + 110100, 76974, 127405, 66340, 59717, 74598, 77, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // OU 公钥 invG(base-2^17,小端序,256 limb) + static const uint64_t ou_invG[ARR_LEN] = { + 75015, 51305, 11174, 79166, 10654, 117345, 79575, 125483, 192, + 111361, 15109, 37776, 85584, 122681, 15179, 105154, 25327, 22587, + 10611, 51779, 65631, 89124, 90845, 91406, 73570, 42237, 118822, + 48197, 100817, 19642, 69889, 42872, 74326, 9921, 113482, 126650, + 61799, 19464, 21459, 1326, 73665, 97400, 6278, 130202, 98895, + 85755, 98021, 67243, 104269, 47402, 94531, 35379, 69858, 42485, + 14950, 53795, 4348, 91203, 443, 82765, 110226, 47814, 97615, + 50660, 105785, 92942, 62675, 126889, 87835, 57079, 60415, 7170, + 37774, 107424, 15927, 51745, 24672, 25017, 106170, 59266, 1628, + 2553, 28098, 44213, 40053, 74899, 36703, 98433, 35250, 63975, + 78378, 39773, 5373, 52527, 114195, 74552, 13400, 119066, 66763, + 69413, 53071, 58437, 15910, 76712, 13865, 10931, 116988, 77261, + 71814, 100804, 100791, 26120, 100684, 120510, 6659, 110701, 71849, + 42678, 3633, 76053, 90059, 80147, 120136, 75883, 64838, 62488, + 96341, 20519, 128162, 92624, 17169, 101609, 30200, 108249, 82915, + 76596, 30678, 114976, 2916, 2981, 56123, 81852, 24638, 52981, + 78410, 117826, 22676, 81030, 33369, 37285, 22870, 37443, 105277, + 20078, 73335, 84923, 33748, 100193, 36046, 40084, 118991, 108731, + 21274, 40120, 129472, 60601, 125338, 18250, 21635, 34396, 126171, + 107271, 62218, 101397, 109320, 91140, 78610, 109963, 44212, 26669, + 123743, 36240, 85404, 52594, 25011, 83228, 46488, 60262, 42699, + 67980, 77026, 105558, 21953, 103173, 41149, 50717, 92954, 61052, + 77499, 68011, 113323, 30031, 48934, 71279, 10577, 16304, 75764, + 104446, 54574, 46616, 108377, 46414, 44739, 26464, 84994, 111443, + 61901, 106984, 103375, 74968, 53441, 83176, 26611, 90857, 85135, + 124470, 116660, 10803, 47700, 73153, 107123, 119153, 23110, 107233, + 119201, 8196, 115181, 16085, 20983, 105180, 27, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // ── Step 1:CPU 端建立预计算表,上传 GPU + // ────────────────────────────────── + const size_t table_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + uint64_t *h_G_table = (uint64_t *)malloc(table_bytes); + uint64_t *h_invG_table = (uint64_t *)malloc(table_bytes); + generate_G_table(Modn, ou_G, h_G_table); + generate_invG_table(Modn, ou_invG, h_invG_table); + + uint64_t *d_G_table = nullptr, *d_invG_table = nullptr; + CUDA_CHECK(cudaMalloc(&d_G_table, table_bytes)); + CUDA_CHECK(cudaMalloc(&d_invG_table, table_bytes)); + CUDA_CHECK( + cudaMemcpy(d_G_table, h_G_table, table_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_invG_table, h_invG_table, table_bytes, + cudaMemcpyHostToDevice)); + free(h_G_table); + free(h_invG_table); + + // ── Step 2:随机生成密文和明文 + // ───────────────────────────────────────────── + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)BATCH * OU_EXP_LIMBS * sizeof(uint64_t); + + // 明文上界(同 ou_mulplain):k[21] ∈ [2^18, P_TOP-1],保证 k < p 且 ≥ + // 2^1362 + const uint64_t P_TOP = 485219ULL; + const uint64_t K_TOP_MIN = 262144ULL; // 2^18 + +#define RAND64() \ + (((uint64_t)(rand() & 0x7FFF) << 49) | ((uint64_t)(rand() & 0x7FFF) << 34) | \ + ((uint64_t)(rand() & 0x7FFF) << 19) | ((uint64_t)(rand() & 0x7FFF) << 4) | \ + ((uint64_t)(rand() & 0x000F))) + const uint64_t MASK17_T = (1ULL << BASE_BITS) - 1ULL; + const uint64_t CT_TOP = Modn[240]; // = 247,密文最高 limb 严格上界 + + uint64_t *h_ciphers = (uint64_t *)malloc(ct_bytes); + uint64_t *h_plains = (uint64_t *)malloc(exp_bytes); + uint64_t *h_results_add = (uint64_t *)malloc(ct_bytes); + uint64_t *h_results_sub = (uint64_t *)malloc(ct_bytes); + + srand((unsigned int)time(nullptr)); + + // 密文:limb[0..239] 随机 17-bit;limb[240] < CT_TOP;limb[241..255] = 0 + for (int b = 0; b < BATCH; b++) { + uint64_t *c = h_ciphers + (size_t)b * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < 240) + c[j] = (((uint64_t)(uint32_t)rand()) ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + else if (j == 240) + c[j] = (uint64_t)rand() % CT_TOP; + else + c[j] = 0ULL; + } + } + + // 明文:22 limb base-2^64;k[0..20] 随机;k[21] ∈ [K_TOP_MIN, P_TOP-1] + for (int b = 0; b < BATCH; b++) { + uint64_t *k = h_plains + (size_t)b * OU_EXP_LIMBS; + for (int j = 0; j < OU_EXP_LIMBS - 1; j++) k[j] = RAND64(); + k[OU_EXP_LIMBS - 1] = K_TOP_MIN + RAND64() % (P_TOP - K_TOP_MIN); + } + + // ── Step 3:设备端分配,H2D 计时 ───────────────────────────────────────── + // d_ct_add / d_ct_sub 各持一份密文(ou_addAndsubp 原地覆写,不能共用) + uint64_t *d_ct_add = nullptr, *d_ct_sub = nullptr, *d_pt = nullptr; + CUDA_CHECK(cudaMalloc(&d_ct_add, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_sub, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_pt, exp_bytes)); + + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(d_ct_add, h_ciphers, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_ct_sub, h_ciphers, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_pt, h_plains, exp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 4:密文加明文(is_add=true) ──────────────────────────────────── + AddSubpTiming t_add = ou_addAndsubp( + true, BATCH, d_ct_add, d_pt, d_G_table, d_invG_table, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 5:密文减明文(is_add=false) ─────────────────────────────────── + AddSubpTiming t_sub = ou_addAndsubp( + false, BATCH, d_ct_sub, d_pt, d_G_table, d_invG_table, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 6:D2H 两份结果,计时 + // ──────────────────────────────────────────── + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_results_add, d_ct_add, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK( + cudaMemcpy(h_results_sub, d_ct_sub, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 7:打印计时 + // ────────────────────────────────────────────────────── + printf("\n========== ou_addAndsubp test ==========\n"); + printf(" batch=%d ARR_LEN=%d OU_TAU=%d TABLE_SIZE=%d\n", BATCH, ARR_LEN, + OU_TAU, TABLE_SIZE); + printf(" [H2D 密文×2 + 明文] : %8.4f us (%6.4f us/op)\n", + ms_h2d * 1000.f, ms_h2d * 1000.f / BATCH); + printf(" [ADD] Getgp kernel : %8.4f us (%6.4f us/op)\n", + t_add.getxp_ms * 1000.f, t_add.getxp_ms * 1000.f / BATCH); + printf(" [ADD] FMLM kernel : %8.4f us (%6.4f us/op)\n", + t_add.fmlm_ms * 1000.f, t_add.fmlm_ms * 1000.f / BATCH); + printf(" [ADD] total : %8.4f us (%6.4f us/op)\n", + t_add.total_ms * 1000.f, t_add.total_ms * 1000.f / BATCH); + printf(" [SUB] Getinvgp kern : %8.4f us (%6.4f us/op)\n", + t_sub.getxp_ms * 1000.f, t_sub.getxp_ms * 1000.f / BATCH); + printf(" [SUB] FMLM kernel : %8.4f us (%6.4f us/op)\n", + t_sub.fmlm_ms * 1000.f, t_sub.fmlm_ms * 1000.f / BATCH); + printf(" [SUB] total : %8.4f us (%6.4f us/op)\n", + t_sub.total_ms * 1000.f, t_sub.total_ms * 1000.f / BATCH); + printf(" [D2H 结果×2] : %8.4f us (%6.4f us/op)\n", + ms_d2h * 1000.f, ms_d2h * 1000.f / BATCH); + printf("==========================================\n\n"); + /* + // ── Step 8:写入测试数据文件 ───────────────────────────────────────────── + // 格式 ou_addAndsubp_test.txt: + // 行 1 (注释):元信息 + // 行 2 (N:) :OU 模数 n(256 limb,base-2^17) + // 行 3 (G:) :公钥 G (256 limb,base-2^17,标准域) + // 行 4 (INVG:):公钥 invG(256 limb,base-2^17,标准域) + // ADD 块:每用例 3 行 ADD_CT / ADD_PT / ADD_R + // SUB 块:每用例 3 行 SUB_CT / SUB_PT / SUB_R + FILE *fp = fopen("ou_addAndsubp_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_addAndsubp_test.txt\n"); + } else { + fprintf(fp, + "# ou_addAndsubp test BATCH=%d ARR_LEN=%d BASE_BITS=%d" + " OU_EXP_LIMBS=%d TABLE_SIZE=%d CT_BASE=2^17 PT_BASE=2^64\n", + BATCH, ARR_LEN, BASE_BITS, OU_EXP_LIMBS, TABLE_SIZE); + // N + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + // G + fprintf(fp, "G:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_G[j]); + fprintf(fp, "\n"); + // INVG + fprintf(fp, "INVG:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_invG[j]); + fprintf(fp, "\n"); + // ADD 用例 + for (int b = 0; b < BATCH; b++) { + const uint64_t *ct = h_ciphers + (size_t)b * ARR_LEN; + const uint64_t *pt = h_plains + (size_t)b * OU_EXP_LIMBS; + const uint64_t *r = h_results_add + (size_t)b * ARR_LEN; + fprintf(fp, "ADD_CT:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", + (unsigned long long)ct[j]); fprintf(fp, "\nADD_PT:"); for (int j = 0; j < + OU_EXP_LIMBS; j++) fprintf(fp, " %llu", (unsigned long long)pt[j]); + fprintf(fp, "\nADD_R:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", + (unsigned long long)r[j]); fprintf(fp, "\n"); + } + // SUB 用例 + for (int b = 0; b < BATCH; b++) { + const uint64_t *ct = h_ciphers + (size_t)b * ARR_LEN; + const uint64_t *pt = h_plains + (size_t)b * OU_EXP_LIMBS; + const uint64_t *r = h_results_sub + (size_t)b * ARR_LEN; + fprintf(fp, "SUB_CT:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", + (unsigned long long)ct[j]); fprintf(fp, "\nSUB_PT:"); for (int j = 0; j < + OU_EXP_LIMBS; j++) fprintf(fp, " %llu", (unsigned long long)pt[j]); + fprintf(fp, "\nSUB_R:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", + (unsigned long long)r[j]); fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_addAndsubp test] 测试数据已写入 ou_addAndsubp_test.txt\n"); + } + */ + // ── Step 8:专项吞吐量测试(多轮平均,消除单次抖动)────────────────────── + // 原理:单次计时受 GPU 冷启动、时钟 boost 未稳、OS 调度抖动影响, + // 直接多轮测量再取平均才是工业标准的 GPU 吞吐量基准方法。 + // 注意:ou_addAndsubp 对 d_ct_add/sub 原地覆写; + // 对全位宽明文(~1363 bit),每轮耗时与输入数据值无关, + // 无需在轮次间重置密文,可直接累加计时。 + { + const int WARMUP_ITERS = 3; // 预热轮数(丢弃,仅稳定 GPU 时钟) + const int BENCH_ITERS = 3; // 正式计时轮数 + + // ── 预热 ────────────────────────────────────────────────────────── + for (int i = 0; i < WARMUP_ITERS; i++) { + ou_addAndsubp(true, BATCH, d_ct_add, d_pt, d_G_table, d_invG_table, + d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + ou_addAndsubp(false, BATCH, d_ct_sub, d_pt, d_G_table, d_invG_table, + d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + } + + // ── 正式多轮计时 ────────────────────────────────────────────────── + double sum_add_ms = 0.0, sum_sub_ms = 0.0; + for (int i = 0; i < BENCH_ITERS; i++) { + AddSubpTiming ta = ou_addAndsubp( + true, BATCH, d_ct_add, d_pt, d_G_table, d_invG_table, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + AddSubpTiming ts = ou_addAndsubp( + false, BATCH, d_ct_sub, d_pt, d_G_table, d_invG_table, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + sum_add_ms += (double)ta.total_ms; + sum_sub_ms += (double)ts.total_ms; + } + + // ── 统计 ────────────────────────────────────────────────────────── + // 总操作数 = BENCH_ITERS × BATCH;时间单位 ms → s 除以 1000 + const double total_add_ops = (double)BENCH_ITERS * BATCH; + const double total_sub_ops = (double)BENCH_ITERS * BATCH; + const double add_tput_ops = total_add_ops / (sum_add_ms * 1.0e-3); + const double sub_tput_ops = total_sub_ops / (sum_sub_ms * 1.0e-3); + const double avg_add_us = sum_add_ms * 1.0e3 / total_add_ops; // us/op + const double avg_sub_us = sum_sub_ms * 1.0e3 / total_sub_ops; + + printf("===== 吞吐量专项测试(预热%d轮 + 计时%d轮 × batch=%d)=====\n", + WARMUP_ITERS, BENCH_ITERS, BATCH); + printf(" [ADD] 平均延迟 : %8.4f us/op\n", avg_add_us); + printf(" [ADD] 吞吐量 : %14.2f ops/s (%10.4f Kops/s)\n", + add_tput_ops, add_tput_ops / 1000.0); + printf(" [SUB] 平均延迟 : %8.4f us/op\n", avg_sub_us); + printf(" [SUB] 吞吐量 : %14.2f ops/s (%10.4f Kops/s)\n", + sub_tput_ops, sub_tput_ops / 1000.0); + printf("======================================================\n\n"); + } + // ── 释放 + // ────────────────────────────────────────────────────────────────── +#undef RAND64 + free(h_ciphers); + free(h_plains); + free(h_results_add); + free(h_results_sub); + cudaFree(d_ct_add); + cudaFree(d_ct_sub); + cudaFree(d_pt); + cudaFree(d_G_table); + cudaFree(d_invG_table); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_addhomo.cu b/heu/library/algorithms/ou_new/ou_addhomo.cu new file mode 100644 index 0000000..0e46704 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_addhomo.cu @@ -0,0 +1,18589 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 11687101653036, 18446743758138104250, 14211471509222940197, + 4235272491137479070, 18240298465959804861, 3612012136679214338, + 17250142769850486366, 16237778209286303069, 4059562170797630074, + 9869391574692501418, 11541833079389403685, 6455235399891900791, + 12385636973677184651, 16642176613718519917, 18371701685023660666, + 12908184103137142305, 432025558415355680, 13421505415041197067, + 3981525027614746084, 11947396656646251748, 17803480186391939853, + 8964175857302805981, 10755671907831908778, 1904820253353466484, + 10527642910445979147, 13113465356235953778, 7234975774234226545, + 9757071334591424272, 7316737729072523157, 6883729646458586186, + 8932421558757596044, 14597307983516136610, 14233691893978885667, + 11827243429966767869, 4500161214914298038, 10200789522299270510, + 13311736198424319672, 15003914107128413955, 15970198162388647651, + 5178000660069144055, 12257227048244175003, 199311466455739912, + 1199638074941611369, 7733792994982443122, 5589666049742506788, + 13804186403958915190, 17844357954872068228, 1608648031291287388, + 14833453796159433687, 17508515457688533059, 5757945642895465137, + 13081582882382389324, 5394527028006066918, 276092650195297071, + 17347510335080686092, 1701269284563161833, 13303804442297711418, + 8253121806998843455, 1803714610749342533, 15051344875346831329, + 17235261944528818002, 8632334691816709691, 7259437303239191782, + 15692170915486480673, 16909097158836466193, 3813100579756643340, + 8120672335331207645, 17658082942068012540, 14712625527555672008, + 3013490167685507391, 17053781224112072993, 7951833678156564288, + 16882129623470333590, 16598623833219974769, 8844055318436626140, + 7109452390029814739, 17973202994978822509, 17988667507327805934, + 4123463954687411915, 4940157814071223716, 10034667236662981698, + 2823422309362454166, 8801094461519448259, 16352417325266577247, + 4120137837347666316, 13914619055383639681, 13924524887541743843, + 16792129803426718801, 17665685411675475587, 4424050859418317508, + 7334727500762095896, 17503735803347907081, 17481354675564559188, + 3960789173235012097, 2382836270529746228, 4344431928835757414, + 10028103253977297380, 10742172769839739841, 4046021704283408275, + 15048763124930553406, 12884462872804057858, 6406919215710243952, + 2305936473049040111, 13121402678735537817, 4908133312471037493, + 11377681924055600462, 12660562847764485932, 4317962794267615766, + 1681209879049333151, 17555975506207345043, 5125753427319466102, + 3335447880892075219, 8356915887911857374, 8584812879987417990, + 17895049134452250854, 9085212301170466888, 7267866673654176139, + 11010976396069557933, 3608178248288276404, 676841772753514930, + 14867830803115014463, 1834280874555657868, 16310358636794094835, + 7330665673582977596, 15741791681143414831, 210676699798228406, + 17198551982049102727, 6109711417879930926, 10546103406410231004, + 15078747006884146927, 15249364729241398859, 7659463200845052268, + 6442795927427660859, 2250605931405808055, 8092318475578226474, + 18259756362431830946, 17518863421902388405, 10430337473484617894, + 11467857499548239759, 14850024957617392139, 4520997378651548347, + 1002071158803202119, 13705130616563662128, 16248739479905188442, + 3181222670542328118, 10465632553806001024, 13994891411389346011, + 10984874395183979750, 9226503659276477277, 16804196055515371214, + 10159135197864231902, 2843438330111370559, 10398977526462394360, + 13629959477406403811, 17564676539751491213, 9169917922917002660, + 9102739415085474595, 2571556159645270064, 15524480249380336561, + 16515752187799014418, 15631770625703612855, 15987278742587054296, + 6287084574542908707, 5110496278831165043, 3153708236541173992, + 12927769102029794814, 15247894309294568262, 9307521059673973398, + 2770646931232435574, 16454025378125412919, 10633343977751814089, + 1784373777963240394, 6934373261888962825, 17349218505508302052, + 10968015619489026254, 1752595997967202555, 12837659792768865327, + 14742784040235182007, 10962127976630829136, 5900158086628228693, + 17316277940499271723, 348494413055888541, 488358709608098659, + 10382164829144562111, 10796492275262689178, 10517686393808888673, + 779545184377220173, 3381663347793212032, 15236282164919686367, + 6334307549334192514, 3063003522052060686, 12114810039050492787, + 8870556826759400006, 4038453701720268007, 14314379071943608840, + 4339980657355467038, 9171890896160995321, 7917821014549284449, + 13180571956635383946, 18248798750186441796, 14577763713235528379, + 1809799949882400184, 6318379507872769143, 1709904138639204836, + 4433595655137503557, 14198791200961044293, 4650959752702308354, + 3318171583106947229, 209694021303727738, 1076839995989163599, + 16905707527716642672, 3695319746732430913, 10252674333030518371, + 16420700828256874086, 2433845634309674122, 16595131843996099370, + 16829576163922136336, 4841410023332049560, 12592434475652961886, + 17572224200405592702, 2431385938003789647, 10061028979483934059, + 2925075822122586101, 8606434114160323337, 5607119374066730577, + 11884780541782053893, 3126131661631420656, 10027052590555524485, + 16048091853305621015, 15852396680435263215, 17385109108245871297, + 12899005442699559936, 7848549015331456901, 9096729734807002481, + 4996467004929486051, 12243245730936161727, 9057745574396783344, + 7771655314147204603, 17881871823990369590, 15212325419875966733, + 11042754214829301385, 3281380445998958437, 17088850123667971831, + 14495632125498788619, 17994272834268936450, 16829150404837372, + 6402480299793703941, 10533393325012975763, 2416806924432625423, + 2875142845742952402, 12269921477737603466, 712826757316093295, + 4415075740707273176, 7975119161839365838, 11673666813117204770, + 9840315584201792241}; + +const uint64_t con_modn_shoup[256] = { + 14252880640204352951, 18322338132012530560, 6849211572557864124, + 15932334426407894751, 14177538223384229049, 9106170071204927292, + 16758578487033067404, 5864117390783715525, 16494545899140440021, + 5897258617717822505, 1933174352700629569, 12810083258791009448, + 12690514865985841899, 3970720354745169798, 4239814533413767079, + 5609102486863112468, 11230284723594595426, 17034417004294615591, + 5986132948557996967, 1868566874544028188, 2158239585541928173, + 17841097290863850509, 9305374060203424222, 10083694270531949160, + 2654649954734757684, 17721823101598855865, 599980504891318557, + 17732121566266018424, 1832524248816892725, 5295674783104160211, + 13283213776815003617, 10900691717351424196, 6021057974446928650, + 12624795618036119368, 2798162278377969124, 7399862538297940480, + 576839721297220460, 16060704215571397670, 16380205270440154947, + 7499979448419237887, 13254841413012858481, 10664669973443882596, + 1765312882701300737, 9426266066319221551, 12823009007753704482, + 10630699336822577567, 16298120910453621338, 13674950148586695572, + 17678273253225120972, 6798806775207395115, 13410427498759750653, + 2614783784077562964, 9342414901595102647, 14373786281595714575, + 6330183866169305354, 998675938268783033, 10221732541776071598, + 8979762078911881033, 9878667596621300249, 4856285279936479033, + 14833025980849776526, 5604878655902262946, 13088552648421720803, + 1801154013367199521, 4158119999782205988, 7904891504652660262, + 9042945763108429841, 4642264688478488779, 16204979912313018920, + 3580705517878336362, 9712754433060271621, 8675179099278674516, + 8186655178929728093, 1884659203003867161, 17775229374523385263, + 6390348527000753038, 6439058351892770174, 2339745637453507323, + 16274314407512660647, 2247518004490005028, 18003796786185432156, + 4540940947376355923, 11987538437574474975, 16166798012420960901, + 16121611900792272328, 16082928115738878740, 121528093926685229, + 11609994994605905995, 2593441955413327993, 16920803883743198476, + 11945409668615507125, 15459882499135165139, 4709903422099897132, + 1915945056478813527, 17487099108173624447, 4121351438621439846, + 11648490996845515622, 16906896413860707859, 440932069689474224, + 5596373384320545758, 6286719224840257488, 15070666469307485122, + 1718056780076659255, 16292491121877970301, 16399246121914763003, + 868264559834958645, 5880650461523548368, 13037697811177873232, + 12598280349103069353, 8787439026840426841, 75102682845531848, + 10793543124523682506, 4058772666671965704, 4575391113880810276, + 7977675084418792789, 3637051392050280908, 16362683407568863478, + 18347383388798481077, 9115743514592553391, 9421569851894468249, + 15594101773322529942, 11807267208082355523, 2211845703086696074, + 17348335706114235958, 11847926503719721254, 17547278040398801999, + 5869056178242350580, 13320003773654588467, 12699478824066277151, + 5239882100070320334, 7261256595809982529, 18328110662323898796, + 14262563528151153433, 10570694000294503318, 12813828885908106934, + 10929763919809758798, 2938308820079533814, 12010181661893483546, + 5724066348617601097, 14693589767406584560, 7346156359909105313, + 12463844683080763521, 6157213913132141689, 10056871538135474507, + 10533920527198962255, 8235152268042627670, 12319030087627737531, + 8540756947882872704, 7431325835456426550, 1301653294393655025, + 5378735902526386334, 14912383613060771114, 639721130790109100, + 16570337161183830109, 3674985562098081097, 25882515425358888, + 5781372063417524209, 16445334884700810773, 3544553957273777187, + 7670642182104980993, 9626872654279745485, 16105190074590295359, + 7770490841006776992, 15228876210409060859, 9849662374993193055, + 8654391918266496929, 6489400423788825431, 3633531112627186925, + 14858949636521671041, 1105232854426717343, 16217154593325743207, + 6106199821319202195, 15821125396653981131, 397434568115144159, + 5468761408652955936, 1296217405573136392, 8004677824854586354, + 2227606275875858042, 810603102045699119, 10604814613007946604, + 5290458938805336352, 17851909192937068431, 13268718299834334195, + 10806219279687765004, 7326952643401977865, 5984244256621982617, + 1659224770285885321, 10142490388121661620, 543184966312697096, + 6161334213393132777, 8606758596526178508, 16789552215120228336, + 8023727110355433065, 4647377966945102884, 6753109868947696463, + 3601294586865406454, 17607120286020078979, 2828322973879047780, + 2912791741152772408, 12563760677479960004, 14534280132691080513, + 1458078075305867271, 965960906590360464, 7718560025107401880, + 5982496231126478023, 9871084626240060187, 13176440103626612310, + 12254705932735020033, 6091017959275360856, 8575903195682037371, + 12248661790925770989, 15874428453561837902, 11580211822667122432, + 1675684581791044528, 953808119473653258, 5212502010992923248, + 11653707338854681558, 2814312282799796103, 10741977896772006037, + 792456711582471131, 5829394638712031862, 14582582386261703791, + 15116195068952376181, 16192690152961596921, 14996186982344279167, + 12452747948715197047, 9822110408961686525, 13084672463213903010, + 15511412028972873361, 4378034570765898628, 7434337709561426193, + 10855303285220731112, 7759227917166922418, 4976939851858003292, + 5204453497818198107, 16768838491807377833, 9561958848337334713, + 703798444570847189, 14816217796224051499, 2968875028840941609, + 6519175664834707328, 17450997194073458375, 12811758118865484166, + 10759621678827047865, 11859099579116844022, 14425180740705110568, + 7511257586748720299, 3539736587530686592, 4447312216206097890, + 11184913710261542236, 8771977661991653642, 12354338902272316516, + 15541265547581959679, 14587017360825887711, 15248329331280087903, + 10992628732632760992}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +int main() { + uint64_t Modn[256] = { + 43351, 84159, 17007, 126963, 115814, 64975, 14865, 122878, 58093, + 76773, 16638, 87086, 110462, 105466, 35053, 36095, 8051, 116177, + 119699, 118157, 12357, 71314, 68424, 35266, 58013, 63468, 22117, + 10903, 124058, 90359, 68490, 117774, 56449, 45990, 26837, 86153, + 120741, 31603, 78596, 24019, 45134, 33649, 61458, 59406, 88868, + 60745, 113313, 123484, 30017, 98185, 93108, 73040, 39521, 18181, + 2647, 51647, 10194, 73702, 22934, 64, 29664, 94536, 9414, + 63827, 6028, 107137, 71399, 49216, 8196, 46100, 117329, 67195, + 25041, 122567, 110161, 82524, 85064, 85420, 38367, 90728, 6216, + 87366, 124652, 29067, 100922, 38894, 64688, 22860, 83774, 130371, + 39036, 94816, 45277, 76221, 67984, 78245, 70889, 64430, 52640, + 50933, 54580, 32496, 95587, 110988, 102834, 68631, 42744, 111149, + 127114, 116295, 108662, 4710, 31837, 15424, 50234, 99229, 61393, + 81585, 33195, 14128, 9168, 55047, 119038, 97329, 43164, 111637, + 39396, 13009, 90209, 92184, 81272, 101938, 57149, 82121, 100630, + 37780, 7881, 13181, 8505, 125111, 43862, 119168, 19431, 80034, + 114187, 71294, 52911, 81495, 14533, 87246, 126978, 30310, 9978, + 44551, 60081, 126942, 75376, 77030, 36034, 104993, 58885, 90371, + 111023, 45378, 97203, 126393, 72942, 8192, 124336, 37338, 116797, + 66693, 60337, 12040, 90738, 108119, 66171, 78981, 79494, 91989, + 89494, 118041, 29798, 30883, 110522, 122729, 7823, 62523, 20666, + 52089, 43045, 51146, 24317, 38753, 122735, 100047, 56716, 117911, + 60032, 27220, 44093, 56, 113553, 49629, 84418, 64845, 67097, + 10050, 8296, 90055, 75973, 63190, 82919, 56713, 30800, 90227, + 63208, 39501, 61899, 129744, 78395, 58460, 121961, 72489, 20054, + 64673, 102069, 663, 109348, 36701, 18676, 98666, 77634, 108307, + 103092, 49932, 49372, 39729, 72337, 90146, 247, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49877, 103520, 14633, 105196, 88121, 75472, 97344, 77893, 104350, + 114542, 91892, 50952, 64070, 21857, 78639, 2966, 6840, 35973, + 24928, 25085, 49586, 41185, 13904, 75536, 2710, 53110, 104559, + 109441, 40809, 28751, 119672, 47293, 31768, 78934, 55525, 52600, + 69445, 27490, 96611, 84604, 11133, 106316, 35838, 34416, 127179, + 34651, 118161, 16314, 45331, 107582, 24235, 12592, 81941, 19278, + 85161, 104589, 13239, 111883, 81763, 64318, 45287, 80867, 37986, + 126633, 118333, 32371, 49588, 111307, 6225, 98536, 50417, 128492, + 7707, 81120, 51703, 13989, 14502, 86184, 74055, 95503, 80070, + 61934, 73173, 322, 128915, 100622, 92460, 3170, 90102, 44305, + 79192, 110628, 84896, 92155, 13698, 129632, 82311, 123141, 99527, + 28216, 52443, 78019, 37062, 11803, 15622, 2177, 42554, 81945, + 37634, 97471, 11261, 30170, 62893, 129991, 77778, 123677, 75667, + 22518, 67693, 79122, 34634, 117067, 33393, 41664, 46880, 63508, + 1202, 120484, 35970, 14884, 64323, 87199, 61041, 17700, 6496, + 128561, 104072, 129267, 31935, 119451, 54205, 55539, 59285, 49953, + 34036, 45769, 28494, 30664, 79616, 47756, 57833, 99317, 122074, + 130764, 65477, 42562, 43963, 97923, 114657, 34735, 19193, 84201, + 24821, 98149, 101368, 100860, 110862, 96240, 101375, 55675, 99994, + 12323, 56026, 120364, 84030, 10327, 108568, 4795, 122128, 3775, + 1740, 113334, 5740, 61052, 15255, 44939, 84950, 4631, 87197, + 63464, 45041, 49844, 102052, 41710, 76594, 96748, 2213, 15033, + 56862, 42121, 22702, 29104, 55955, 32193, 62378, 61812, 37549, + 27929, 118796, 116386, 35884, 83278, 116744, 103768, 106752, 29801, + 6976, 81713, 55669, 12038, 51733, 6915, 128541, 82038, 102167, + 64630, 125581, 69829, 79662, 80895, 89416, 41571, 113918, 73736, + 22655, 72892, 97009, 75512, 83469, 50798, 35893, 72631, 27752, + 114176, 116066, 35199, 11556, 117400, 53979, 71662, 76589, 35790, + 51797, 38295, 48839, 44050}; + uint64_t R1[256] = { + 95787, 87855, 74590, 120089, 96462, 125333, 89873, 62820, 56744, + 93675, 114260, 58407, 55044, 4742, 20922, 129032, 18634, 103071, + 2852, 114517, 116272, 79216, 95365, 35177, 128432, 80425, 18923, + 592, 13588, 42144, 48019, 39668, 66805, 42663, 33194, 65911, + 93428, 9610, 76041, 48300, 121686, 67062, 30099, 86626, 99273, + 90908, 66468, 28073, 71719, 43868, 40340, 88274, 109318, 21824, + 16472, 116161, 1346, 106033, 20342, 35258, 20632, 105594, 118266, + 97653, 97643, 7306, 67863, 79950, 112151, 117205, 39906, 100559, + 19328, 14826, 43881, 64539, 123341, 37113, 15909, 43631, 27755, + 12868, 87791, 110907, 2763, 41576, 76238, 21079, 28709, 67173, + 22692, 45867, 80137, 111080, 90017, 19215, 15056, 7843, 34411, + 10397, 47157, 113197, 77959, 43337, 123310, 26898, 103324, 95568, + 37773, 53742, 58532, 64900, 12429, 109482, 75505, 70429, 89935, + 67404, 103144, 45403, 28839, 100826, 80183, 60279, 60825, 67114, + 15456, 95163, 5820, 106812, 38605, 127798, 43023, 23037, 109334, + 82354, 36764, 29882, 1460, 109709, 70002, 61938, 129339, 12574, + 23578, 116415, 67219, 51854, 56951, 3482, 98043, 101818, 116934, + 91679, 101444, 12152, 24202, 116763, 23200, 65030, 72892, 11401, + 89306, 104758, 94473, 69024, 2331, 120575, 38998, 71116, 123520, + 12246, 31511, 15417, 54824, 35449, 51579, 129451, 119392, 117118, + 124644, 45532, 83498, 101978, 4264, 7165, 33903, 43033, 21694, + 89206, 4834, 104846, 13789, 111463, 68368, 35303, 42736, 93343, + 26110, 67668, 100827, 18175, 80987, 95026, 43049, 29875, 73340, + 103134, 127164, 12263, 68500, 44986, 40419, 118430, 83125, 27361, + 83899, 599, 102652, 96223, 98631, 95859, 60979, 40494, 118153, + 71598, 113945, 125768, 2378, 115813, 13443, 7038, 3598, 53377, + 26625, 110573, 3294, 119847, 72218, 85832, 125, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 2023, 29313, 55760, 30415, 117209, 68859, 126088, 105595, 44062, + 129614, 108803, 62812, 119884, 130526, 4554, 108683, 53422, 114202, + 34719, 66446, 65067, 23670, 9591, 82680, 20040, 77609, 34903, + 85401, 37760, 43899, 14395, 67595, 12710, 101806, 77200, 103316, + 65099, 125953, 105217, 67808, 5979, 63677, 12355, 111674, 70131, + 45594, 117232, 10425, 81356, 112691, 14758, 10490, 24556, 27922, + 32980, 20928, 118420, 17204, 4244, 126937, 116849, 106497, 51321, + 114935, 45503, 15461, 59271, 111583, 30113, 103352, 10622, 32510, + 41116, 86928, 10137, 101567, 30707, 124356, 108755, 8203, 11158, + 9603, 114740, 5093, 13054, 61800, 75687, 38080, 11550, 87153, + 33247, 18929, 66437, 13511, 39575, 18765, 61155, 77315, 112366, + 76906, 125693, 40793, 40582, 43161, 81338, 111531, 84813, 49322, + 71309, 83250, 14948, 44745, 13967, 98243, 116072, 5842, 82567, + 77993, 80649, 107659, 66320, 122438, 54394, 54983, 79006, 105681, + 94840, 79085, 41950, 106863, 130420, 89427, 83726, 86511, 44750, + 12837, 47751, 33678, 115313, 66053, 43941, 100068, 107956, 169, + 62673, 70167, 106884, 28460, 37125, 129538, 93387, 56010, 35558, + 40621, 77463, 39765, 16013, 85203, 39465, 19946, 122865, 58068, + 60861, 54036, 48234, 130529, 59321, 83170, 116672, 11733, 94357, + 35207, 46800, 22992, 47306, 80973, 28136, 59828, 4338, 109019, + 16604, 58136, 123247, 75151, 11948, 105570, 86992, 72951, 15873, + 12088, 27387, 60538, 80333, 72819, 97245, 41421, 126078, 82127, + 78747, 70199, 99633, 38926, 56088, 20821, 130509, 125699, 127559, + 47166, 102354, 71946, 79468, 81253, 11120, 37016, 20384, 95174, + 126298, 33756, 60551, 106592, 43758, 59029, 22997, 63753, 65367, + 47855, 5100, 118402, 102326, 108278, 49481, 99388, 114710, 101475, + 124482, 27406, 78309, 52397, 69291, 787, 10, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ========================================================================= + // ou_addhomo 正确性测试 + // + // 流程: + // 1. 随机生成 BATCH 对密文 c1、c2(base-2^17,值域 [0, n²)) + // 2. H2D 传输 c1、c2,用 cudaEvent 计时 + // 3. 调用 ou_addhomo 执行 FMLM(c1, c2),内核计时由函数返回 + // 4. D2H 传输结果,计时 + // 5. 将 n²、c1、c2、result 写入 ou_addhomo_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 100000; + const size_t batch_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + + // ── Step 1:随机生成 c1、c2 ─────────────────────────────────────────── + // limb[0..239] 随机 17-bit;limb[240] < Modn[240];limb[241..255] = 0 + // 确保生成值 < d_Modn(OU 模数 n) + const uint64_t MASK17_T = (1ULL << BASE_BITS) - 1ULL; + const uint64_t TOP_LIM_T = Modn[240]; // = 247,d_Modn 最高非零 limb 的上界 + + uint64_t *h_c1 = (uint64_t *)malloc(batch_bytes); + uint64_t *h_c2 = (uint64_t *)malloc(batch_bytes); + uint64_t *h_result = (uint64_t *)malloc(batch_bytes); + + srand((unsigned int)time(nullptr)); + for (int b = 0; b < BATCH; b++) { + uint64_t *c1 = h_c1 + (size_t)b * ARR_LEN; + uint64_t *c2 = h_c2 + (size_t)b * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < 240) { + c1[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + c2[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + } else if (j == 240) { + c1[j] = (uint64_t)rand() % TOP_LIM_T; + c2[j] = (uint64_t)rand() % TOP_LIM_T; + } else { + c1[j] = 0ULL; + c2[j] = 0ULL; + } + } + } + + // ── Step 2:设备端分配 ───────────────────────────────────────────────── + uint64_t *d_c1_t = nullptr, *d_c2_t = nullptr; + CUDA_CHECK(cudaMalloc(&d_c1_t, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_t, batch_bytes)); + + // ── Step 3:H2D 传输,计时 ───────────────────────────────────────────── + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(d_c1_t, h_c1, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_c2_t, h_c2, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 4:GPU 同态加法,内核计时由 ou_addhomo 返回 ────────────────── + // 注意:ou_addhomo 原地覆写 d_c1_t,调用后 d_c1_t 存储结果 + float ms_kernel = ou_addhomo( + BATCH, + d_c1_t, // inout:c1 → FMLM(c1,c2) + d_c2_t, // 只读:c2 + d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 5:D2H 传输结果,计时 ───────────────────────────────────────── + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_result, d_c1_t, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 6:打印计时 ─────────────────────────────────────────────────── + printf("\n========== ou_addhomo test ==========\n"); + printf(" batch=%d, ARR_LEN=%d, BASE_BITS=%d\n", BATCH, ARR_LEN, BASE_BITS); + printf(" c1+c2 H2D : %8.2f us (%6.4f us/op)\n", ms_h2d * 1000.f, + ms_h2d * 1000.f / BATCH); + printf(" GPU kernel : %8.2f us (%6.4f us/op)\n", ms_kernel * 1000.f, + ms_kernel * 1000.f / BATCH); + printf(" result D2H : %8.2f us (%6.4f us/op)\n", ms_d2h * 1000.f, + ms_d2h * 1000.f / BATCH); + printf("=====================================\n\n"); + /* + // ── Step 7:写入测试数据文件 ──────────────────────────────────────────── + // 文件格式(ou_addhomo_test.txt): + // 行 1:注释 + 元信息 + // 行 2:n² 的 ARR_LEN 个 limb(空格分隔) + // 此后每个测试用例 3 行:c1 limbs / c2 limbs / result limbs + FILE *fp = fopen("ou_addhomo_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_addhomo_test.txt\n"); + } else { + // 元信息行 + fprintf(fp, "# ou_addhomo test BATCH=%d ARR_LEN=%d BASE_BITS=%d\n", + BATCH, ARR_LEN, BASE_BITS); + // OU 模数 n(= d_Modn)一行 + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + // 每个用例 3 行 + for (int b = 0; b < BATCH; b++) { + const uint64_t *c1 = h_c1 + (size_t)b * ARR_LEN; + const uint64_t *c2 = h_c2 + (size_t)b * ARR_LEN; + const uint64_t *r = h_result + (size_t)b * ARR_LEN; + fprintf(fp, "C1:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c1[j]); + fprintf(fp, "\n"); + fprintf(fp, "C2:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c2[j]); + fprintf(fp, "\n"); + fprintf(fp, "R:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)r[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_addhomo test] 测试数据已写入 ou_addhomo_test.txt\n"); + } + */ + // ── Step 7:专项吞吐量测试(多轮平均,消除单次抖动)────────────────────── + // ou_addhomo 原地覆写 d_c1_t;FMLM 操作耗时与数据值无关, + // 无需在轮次间重置密文,可直接复用已有 d_c1_t / d_c2_t。 + { + const int WARMUP_ITERS = 3; + const int BENCH_ITERS = 10; + + // 预热(稳定 GPU 时钟,丢弃计时) + for (int i = 0; i < WARMUP_ITERS; i++) { + ou_addhomo(BATCH, d_c1_t, d_c2_t, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + } + + // 正式多轮计时 + double sum_kernel_ms = 0.0; + for (int i = 0; i < BENCH_ITERS; i++) { + float t = ou_addhomo( + BATCH, d_c1_t, d_c2_t, d_negModn, d_con_NegModn_shoup, d_Modn, + d_con_Modn_shoup, MOD, d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample); + sum_kernel_ms += (double)t; + } + + const double total_ops = (double)BENCH_ITERS * BATCH; + const double avg_us_op = sum_kernel_ms * 1.0e3 / total_ops; + const double bench_tput = total_ops / (sum_kernel_ms * 1.0e-3); + + printf("===== 吞吐量专项测试(预热%d轮 + 计时%d轮 × batch=%d)=====\n", + WARMUP_ITERS, BENCH_ITERS, BATCH); + printf(" 平均延迟 : %8.4f us/op\n", avg_us_op); + printf(" 吞吐量 : %14.2f ops/s (%10.4f Kops/s)\n", bench_tput, + bench_tput / 1000.0); + printf("======================================================\n\n"); + } + // ── 释放 ─────────────────────────────────────────────────────────────── + free(h_c1); + free(h_c2); + free(h_result); + cudaFree(d_c1_t); + cudaFree(d_c2_t); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_dec.cu b/heu/library/algorithms/ou_new/ou_dec.cu new file mode 100644 index 0000000..9071e90 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_dec.cu @@ -0,0 +1,19721 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12420598007131, 404862574961, 1489476715458592911, + 16957267883830147813, 12043402377674848005, 16493683929612125450, + 6784216004999364667, 1572187082686061104, 12340520880335831418, + 9331149018221532047, 14282361491840612423, 8815005708360508971, + 3015263268809804842, 17865818108905781022, 9920159502644748834, + 16663441635849680526, 16754665591114057436, 18269507414440088358, + 10584069732924195264, 12369982964838682500, 17245146333661260496, + 14241460736918561983, 1495228790312205286, 650596912814140113, + 13561313259753582639, 12825917001300214751, 17648203335955277871, + 4937017841769177676, 8762411814459186696, 10916964928297372446, + 7934650143980081733, 16270299917683438672, 10599916200641161267, + 8153067147237791839, 3263540562937225549, 12427502011302527075, + 233468242765771385, 16905948202777045661, 7718378239207545448, + 12556027611679852532, 14355997036124831232, 10480393087867199263, + 8928529658741694940, 5329501958372121449, 14697840141414380795, + 4109255465868101899, 2189093870568573914, 2015298188623315682, + 15846288464350737625, 6154716281955634074, 14926805960051698119, + 6358221390244429391, 565479903706287841, 544001290941279643, + 8964697110262868308, 18143181760096837871, 4558614659708563008, + 13794757981490580748, 8464822194361423649, 3215767361281822704, + 3162014795756580172, 4046965963624353869, 13045127494050167872, + 2499196209129003472, 9910258488503570389, 17361620693833999278, + 9909125957613821241, 9160076386865281785, 1124681048989805109, + 13488748519780692092, 10062780796386632218, 1229663950817871070, + 9195972272372452954, 16317888560940382210, 7613588099474937282, + 4817436175622690809, 2044827062874516597, 10883598385799752225, + 15983746154563805830, 15305213912578498862, 15429027565390837948, + 15165749666894862326, 17094752132458668719, 17119153326423120334, + 16418978274070441262, 4020071962619695082, 1332089103620365082, + 2185942973674898504, 10623386207974578633, 17221670917682379713, + 1728235680214257322, 4351774997312820994, 13118214034020006923, + 17372706552005974065, 3635082452886395532, 9785476544995844000, + 5239966118419100751, 11446467251832755986, 8441993571518844572, + 17814051721047964692, 6635316643940291166, 591992533123407590, + 5392319256964754787, 6876233685503189066, 12560692609570806257, + 15922352178387947359, 10408865844002002467, 8334121862371121723, + 12325068481620777235, 16272399208182333772, 95913092914333903, + 9337511878088658138, 3620707693508746669, 2410461826666169902, + 15825609924068456517, 14631303029849624577, 12745725307310343054, + 11589389212418519430, 2402535340645687114, 10460197248544579201, + 2150167009874019400, 6579654444057751335, 10629733729325322988, + 15336435747150065712, 2033376745185295929, 5536738554015879544, + 6536459867767524988, 17547258140832520218, 16277632797518007476, + 17414093779241440205, 15172877332690969372, 11510709146567555164, + 4030987138388375538, 4545399202703904429, 3071928236698087396, + 6235823038589723214, 1385517152964008659, 5315971425862540039, + 737213825926788813, 12647942483396120777, 11729199013777796647, + 16701826799686034287, 10264009214043680734, 15712857482319384195, + 9683635541859048778, 12510245784932525987, 11453973523166319383, + 17641407202450751284, 7076031528159713278, 4418210339976245128, + 7042975543133840791, 2687978692681798677, 6620884024635649052, + 18356048427182277532, 12306609112475240772, 9532567709512160247, + 18330814574547253380, 1570356219924936203, 8589904627747755897, + 15801551039736301865, 14430949397822386754, 17457931947828226443, + 17066621315075642576, 5110941585241228080, 3113435720621532772, + 9109366487892874842, 6870609824263132877, 12909242032068451345, + 9978696359698762281, 2334245554206682016, 5474675901673095882, + 6791983755440567441, 7940389551176083495, 8857160764169799470, + 5079795210984134207, 13270967636848123637, 2817746455641616989, + 6823734589021892698, 6695712168139536841, 10011091548170035449, + 1880778261691426396, 4901098797777619034, 10275511334773938844, + 14807980597086261468, 2140936956496274302, 16031412187684387227, + 13312372442173260977, 3698443504872582944, 13391615949904799953, + 9911497060502159398, 6476394081735368123, 10884190273162182630, + 16812094898227487679, 17541511080948430118, 608781942072315334, + 11293244422430215336, 11785770688631267218, 1050397619038471432, + 8886815311739454498, 4768360092769081595, 6853419982646256117, + 4533330592021381799, 4065300292447750493, 5995762058657812513, + 5850499564985481696, 17494639705373259696, 5311542451880315192, + 10494352427384472421, 7053555067647955415, 13249062068286530876, + 4950095133067012336, 978416016041957203, 16132684490165566392, + 9461678038864020399, 13953140033946412303, 3128510702558997798, + 9989688919461798931, 7729804782429579018, 14657677073053593359, + 2033538150495053663, 18287519723785336581, 15945421906051418344, + 827199110667409676, 2705644704721997071, 7966095645914939971, + 16210256519322987223, 13500279664518420546, 6403250311096390047, + 15684661198519501341, 17612977561176850685, 1595673854820173645, + 11355330109781124768, 15152054633977714227, 7594126125750130131, + 1770764465613758168, 16103674937332423683, 12989244865424659639, + 985924609341537427, 17869081554936260320, 13756399484785420256, + 9417665840560873340, 6596619827292141536, 9282779747668579711, + 14087089070266996746, 16797541746528108966, 3649008588426316857, + 305957104569710731, 16307164143397236406, 8758406156783513520, + 8985939822111752351, 9823364479239409710, 14242788910385080985, + 16257804781363675711, 3460875886956782926, 5792587040463753957, + 8058808855494344097}; +const uint64_t con_modn_shoup[256] = { + 17258949916450502562, 8888112940397698219, 713580727021815326, + 7267352695482246207, 16882106641962718850, 4484957681626972705, + 11436540044429239973, 11422524016867871168, 14020805313791955351, + 3223598656225742213, 5973835375502394258, 18389866180389685123, + 16986950504460752843, 14151661301513655720, 17077263121249724353, + 9230648276690066133, 14180947083865490845, 9351212926919619729, + 14132401609129010813, 5094848903632022188, 9701476928191356694, + 2266122304659313991, 9642600553848677189, 17694947392610440444, + 4166068824698342286, 5461002667202655000, 12457224361386750838, + 17183384576469292838, 4577122507392282256, 212990517287292298, + 9593917155195774666, 16026775048324125596, 16942601617377608128, + 14517179166542822181, 5974534123485502056, 3098739027972382032, + 10256410151333955680, 17520563159797137486, 1437139722912674326, + 8992878724252091890, 12528748759234609070, 8869909966108752023, + 921545360850748754, 9072672431128383783, 5738738820990719155, + 14467747599114998535, 17288556119140904797, 10287232211619750677, + 14260418889536810543, 14236170919296794224, 3150044698371016975, + 3734472856964467879, 17367148054045259108, 8704684147659221983, + 2508292278448591736, 8252783960354095393, 5505177864158443900, + 6073691067055527962, 1797426869555915200, 10910566681039151494, + 7942925665011098265, 2975836443917988571, 12168486674970259490, + 16110198076303607574, 3176765343615554343, 16298611667673954183, + 17549972852929401068, 11301222901206006551, 16909965679872285631, + 4718899407226931290, 3249759906455460509, 12902044256019035476, + 2933343404927512513, 12635530155166610452, 1019795786854894305, + 5558319569785986415, 1736084326860448187, 1046201637077468719, + 8212969089240241408, 14337127378294037853, 386901202666220852, + 7007901116396992081, 12993354667310087947, 10362596985035361337, + 2632315220458003599, 15774429557071893885, 16402946062094934668, + 1455058537672538730, 10934206267514764178, 8653495955263034803, + 15626427345777561198, 9949509055664635150, 13551873570797002283, + 8464454269390886841, 17029555676738610993, 2576180826541960979, + 9277239235861014619, 5949389895883929521, 5961129371502370384, + 205015922583922578, 10998263302456627347, 630682958888475996, + 8272932531190805541, 11739326145926911850, 3711218700419121610, + 1995808726167056948, 9541493218673654357, 10947701895562230579, + 13307121742110758926, 5577478578055132785, 17291845771099354096, + 18386024476421024505, 12682871500298847381, 11130873186526484524, + 9511555001775281238, 12573445966195088218, 8735156148113491387, + 11986977738879871749, 3957140364191407682, 9447015894594054903, + 15075232292569908924, 3903378683888054983, 18108762427411217092, + 3848515571709448978, 9812945716174782266, 1439686790983251294, + 16833581092072042070, 16210195608156458739, 5586553670046771142, + 16285469091149528676, 9398773791615283024, 6409294867327805667, + 7451696515003176156, 9352800728836756410, 13006499844621233701, + 8943861815085160245, 9678575573847830426, 8255945162964001220, + 10730739327269562739, 5334149872806451563, 17598838495642584396, + 16369098332122517680, 17955871179034626630, 14641785369093985684, + 14862588269370659671, 18285743176298903762, 13037469746309732440, + 14680321993850845543, 4817891414280274891, 5149080315255804853, + 2128397328110922197, 10589038400063582404, 7361242319334319176, + 2633416174015188179, 8731722470221646364, 17695241939569347999, + 511281303145693578, 3915229765025655898, 12611201828639463913, + 12923773568272819998, 6599022632397876071, 13420679683499938789, + 2049055346114188916, 12969132418156712259, 4309319847942653387, + 18244990947493391329, 4321703802367122695, 6192941661180839271, + 3809220907288466934, 9686397595036380761, 5512567166054515334, + 7602079309929874744, 13912271487139773395, 14184657929002112052, + 14576671517535568717, 10645713420980679390, 6189820458973792135, + 11814662608890796689, 14130844520328466724, 5920079432947567926, + 11563681922574043697, 1737013972124835278, 9154718647461864771, + 4872088247015674989, 3298657241789078374, 4274525812376499266, + 2844739318930852390, 5423787114207982915, 711603148563500498, + 5237712164936082793, 16538115545248024540, 6726637234751909654, + 12914942507138411381, 13794083152765605836, 8092188286693750926, + 3585046709821797374, 9022641750310238763, 465420587122630117, + 13564884715383175347, 11129511873502417680, 5164151854448859979, + 15029903335193380621, 10868746496347154392, 4435651212450171864, + 13217684022763318278, 5933919845722166548, 3356965223531861545, + 5310913073599813054, 15054780849600338193, 6396900692144724485, + 7035542304377335696, 10193072637584812570, 13651162572970216008, + 3375517112149222934, 4439239668198625265, 4272637244444960409, + 12817026244019634782, 69354263255062928, 15302789662880427742, + 347862105615759955, 17769797205451930634, 9642432860176744960, + 12464642972155630307, 15664388251667466866, 11631517894848881873, + 17036907573636512589, 7095039048837509769, 17513417735648162476, + 17481830685724223605, 17280286146669482802, 2842388552256186635, + 3920667164872509047, 12274188700851890415, 10745285319594444214, + 1477783310694359069, 9823805546543362486, 13036068216114647306, + 2630587855422542122, 11133180172685530573, 5498501119950892568, + 12148543347986984583, 8699653759414714944, 12680167660387601338, + 7199244163763404941, 11286265673566327732, 16227882229914399460, + 13238315363757133297, 15029696251503760467, 3493558449506003864, + 1655355875307187670, 2509903151438279534, 16562697969364590150, + 16509371521602835055, 6068654425169380909, 12704430300130691212, + 1645898033077800833}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define OU_T_TAU 256 // 解密指数 t 的 bit 数(p-1 的大素因子,~256 bit) +#define OU_T_EXP_LIMBS (OU_T_TAU / 64) // = 4,每个 t 占 4 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +// 保留参数兼容调用处,128-bit 时不需要 p +static void ou_gen_r_prime(const uint64_t * /*ou_p_limbs17*/, int batch, + uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================================= +// ou_encrypt +// +// 功能:批量 OU 加密(正明文) +// c̃[i] = FMLM(G^{m_i}, H^{r'_i}) mod n (FMLM 域输出) +// +// 数学依据: +// OU 加密公式:c = G^m · H^r mod n +// FMLM 域语义:FMLM(Ã, B̃) = ÷B̃·R⁻¹ mod n +// 因为 Getgp/Getrn 输出 FMLM 域(乘了 R),所以: +// FMLM(G^m·R, H^r'·R) = G^m·H^r'·R mod n ✓ +// +// 三步流程: +// Step① Getgp :G 预计算表 × m_i → G^{m_i} (FMLM 域)→ d_ct_out +// Step② Getrn :H 预计算表 × r'_i → H^{r'_i}(FMLM 域)→ d_Hr(临时) +// Step③ XYfixWarpVector:d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) +// +// 调用前准备: +// generate_G_table(Modn, ou_G, h_G_table) 并上传 → d_G_table +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM +// 域,generate_G_table 输出) d_H_table [TABLE_SIZE × ARR_LEN] H +// 预计算表(FMLM 域,generate_H_table 输出) d_m_batch [batch × +// OU_EXP_LIMBS] 明文 m(base-2^64 小端序,OU_EXP_LIMBS=22) +// d_r_prime_batch [batch × OU_HR_EXP_LIMBS] 随机指数 r'(OU_HR_TAU=128 +// bit,OU_HR_EXP_LIMBS=2) d_ct_out [batch × ARR_LEN] 输出密文(FMLM +// 域) batch 明文数量 +// +// 返回:GPU 端 Step①②③ 总耗时(ms) +// ============================================================================= +float ou_encrypt(const uint64_t *d_G_table, const uint64_t *d_H_table, + const uint64_t *d_m_batch, // [batch × OU_EXP_LIMBS] + const uint64_t *d_r_prime_batch, // [batch × OU_HR_EXP_LIMBS] + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM 域) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getgp — G^{m_i} mod n(FMLM 域)→ d_ct_out + Getgp<<>>( + d_G_table, d_m_batch, OU_TAU, d_ct_out, batch, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step③:XYfixWarpVector — d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) + // 结果:c̃[i] = G^{m_i} · H^{r'_i} · R mod n(FMLM 域密文) + XYfixWarpVector<<>>( + d_ct_out, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================= +// ou_dec +// +// 功能:GPU 批量 OU 解密第一步——模幂 c^t mod n(普通域输出) +// +// OU 完整解密流程: +// Step①(本函数): c^t mod n ← FMLE_mod3_Kernel,NTT 参数为 n 模数 +// Step②(CPU 端): (result) mod p² ← 因 n=p²·q,c^t mod n 再 mod p² = c^t +// mod p² Step③(CPU 端): m = L(c') · gp_inv mod p,其中 L(x) = (x-1)/p +// +// 输入 d_c_tilde 为 FMLM 域密文 c̃ = c·R mod n; +// FMLE_mod3_Kernel 跳过 Step8(输入已是 FMLM 域), +// 保留 Step13(乘 R⁻¹,还原为普通域),输出 c^t mod n。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(FMLM 域,mod n) +// d_t_exp [batch × OU_T_EXP_LIMBS] 指数 t(压缩格式:OU_T_EXP_LIMBS=4 个 +// uint64_t, +// 小端序,位 i 在 limb[i/64] 的第 i%64 +// 位) +// d_output [batch × ARR_LEN] 输出:c^t mod n(普通域,base-2^17) +// batch 批大小 +// d_r0 [ARR_LEN] r₀ = (2^(ARR_LEN×BASE_BITS)-1) mod n +// (FMLM 单位元,即蒙哥马利域中的 1) +// 其余为 NTT 参数(n 模数,与 ou_encrypt 完全一致) +// +// 返回:GPU 端到端耗时(ms) +// ============================================================= +float ou_dec(const uint64_t *d_c_tilde, const uint64_t *d_t_exp, + uint64_t *d_output, int batch, const uint64_t *d_r0, + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * OU_T_EXP_LIMBS; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_t_exp + exp_off, OU_T_TAU, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +int main() { + uint64_t Modn[256] = { + 90459, 46625, 4442, 112182, 114968, 71534, 86507, 127438, 1685, + 52931, 104150, 120548, 110008, 7428, 47921, 17877, 51604, 95338, + 118828, 83898, 128532, 25064, 75321, 17612, 57426, 74847, 91485, + 17341, 32555, 39547, 13559, 24145, 41272, 116851, 19414, 130804, + 122030, 33140, 103122, 68029, 121047, 93139, 39899, 122526, 23084, + 120314, 105769, 34702, 77999, 9851, 111681, 49144, 6962, 87978, + 73280, 9230, 120490, 99393, 36052, 66751, 52136, 26210, 79015, + 53220, 88620, 106517, 78130, 39022, 74644, 83791, 78431, 76879, + 126508, 112379, 112791, 45016, 126862, 98555, 67445, 123550, 78839, + 94946, 54918, 116122, 88926, 128743, 56161, 106652, 89566, 6771, + 65824, 105202, 2315, 108169, 98258, 27489, 25367, 128358, 45812, + 97612, 65111, 13602, 17793, 42341, 58555, 122795, 20331, 7114, + 60301, 88265, 43599, 10109, 18770, 82809, 13834, 117670, 103794, + 49075, 68168, 87021, 105358, 120278, 55703, 7511, 51479, 98172, + 58294, 99843, 55022, 33556, 19683, 122106, 75656, 112874, 46791, + 112145, 106681, 122130, 87140, 32269, 9030, 59222, 342, 75944, + 94103, 130245, 53043, 71941, 64227, 216, 36445, 98379, 21169, + 62160, 91119, 86471, 3598, 23603, 98537, 23157, 124090, 112187, + 67586, 111947, 77293, 12685, 36844, 113775, 50048, 260, 118397, + 88108, 66275, 63944, 42517, 39246, 24220, 115910, 91147, 23953, + 45252, 7345, 75180, 16316, 123347, 38049, 74728, 78637, 21442, + 41317, 90948, 49684, 37151, 40998, 36225, 79231, 81693, 87993, + 55886, 2798, 82414, 110989, 110453, 57672, 14220, 82985, 1843, + 128765, 104271, 125420, 13546, 63807, 90504, 103167, 51426, 33297, + 102375, 75357, 118381, 38152, 11460, 114990, 16009, 63757, 78839, + 46610, 83493, 100684, 103041, 23340, 76918, 82210, 94466, 46728, + 76817, 44869, 84237, 127471, 98069, 60741, 890, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 89416, 26754, 112835, 104869, 63192, 109205, 3104, 29726, 113565, + 56327, 49094, 67751, 77782, 27488, 119631, 71590, 123612, 97756, + 84117, 11401, 55613, 66158, 70476, 101675, 91512, 52471, 103039, + 24137, 47894, 66258, 102515, 73290, 122151, 114889, 128986, 74711, + 74396, 24226, 116902, 16522, 70549, 97827, 20059, 59149, 52686, + 26411, 100945, 78310, 25925, 67014, 79121, 123619, 40826, 66607, + 30092, 89869, 105580, 20910, 93204, 44105, 73778, 21121, 113212, + 113354, 12462, 2188, 4706, 7519, 55812, 128963, 35815, 6339, + 118625, 95153, 80551, 71853, 110048, 116886, 28114, 36427, 99347, + 11391, 127843, 41917, 14312, 58727, 41417, 4720, 129585, 25994, + 106019, 13449, 40561, 107035, 2390, 33535, 7847, 22268, 12977, + 53724, 114048, 27127, 106539, 77244, 114474, 118543, 104524, 97508, + 116569, 73049, 95427, 14163, 96131, 121201, 90712, 20571, 129841, + 128480, 52654, 46075, 61521, 53523, 19041, 42853, 127248, 11120, + 123997, 130413, 56569, 98615, 56998, 99234, 71154, 41850, 65057, + 46995, 104286, 37295, 68580, 38759, 28487, 15348, 68337, 102177, + 78016, 31678, 52663, 1195, 70670, 125162, 16806, 115057, 38759, + 1825, 67883, 105361, 112649, 60917, 88939, 64087, 91926, 11035, + 53857, 46013, 61091, 124646, 128853, 106128, 66902, 108804, 8022, + 25699, 7098, 119054, 99103, 128189, 128398, 1319, 96584, 80228, + 81809, 54282, 6764, 98910, 13622, 81427, 9254, 125751, 67112, + 129121, 60904, 13975, 42521, 71104, 19357, 130027, 15821, 92416, + 81803, 63240, 43145, 90507, 34211, 103714, 62407, 43497, 11360, + 9542, 117006, 82808, 14980, 96632, 5649, 54610, 109476, 15174, + 113385, 48087, 34882, 126953, 67608, 7544, 126322, 75249, 4473, + 4549, 79315, 60258, 82932, 85897, 13892, 119911, 78557, 112723, + 15372, 51919, 45301, 73004, 122707, 121269, 108268, 19504, 49759, + 64648, 50866, 115280, 49268, 9806, 71075, 96540, 114884, 100204, + 79963, 109285, 92516, 38345}; + uint64_t R1[256] = { + 63840, 77643, 38589, 103335, 39211, 83771, 23554, 103602, 13011, + 33265, 99850, 107872, 65932, 9450, 98916, 44944, 70258, 117509, + 29436, 61780, 95269, 57764, 113636, 76719, 51738, 24730, 60149, + 46944, 29753, 55286, 121585, 3994, 70184, 56425, 18607, 115724, + 66839, 44507, 95968, 14278, 79412, 78754, 21044, 58743, 116286, + 99739, 26812, 85232, 130351, 42770, 22378, 65680, 11282, 130626, + 67800, 75341, 43729, 3643, 41124, 90487, 21576, 61026, 35178, + 84381, 65484, 94672, 34715, 116332, 18621, 103526, 117440, 57014, + 7488, 192, 2821, 11532, 82352, 92862, 129569, 103329, 115342, + 91522, 84549, 112318, 128186, 46600, 126196, 12657, 116, 12915, + 15556, 72466, 25692, 116854, 113547, 7800, 66899, 57704, 21718, + 49267, 129570, 121088, 11232, 26664, 76057, 98161, 88201, 110810, + 67711, 57350, 129067, 55009, 66719, 82831, 16394, 45675, 54590, + 128704, 53971, 119083, 90622, 38432, 89295, 6632, 106784, 29141, + 112848, 12094, 121287, 45606, 18268, 123439, 23933, 114383, 91280, + 32929, 75475, 33289, 81297, 101051, 117022, 93617, 78312, 23394, + 80435, 107542, 74032, 67728, 108966, 130175, 112943, 113217, 3110, + 3831, 11764, 90166, 19552, 130175, 21228, 105234, 12469, 127827, + 48250, 34863, 105366, 25850, 91309, 59492, 11089, 37416, 86082, + 43171, 97876, 17564, 124193, 22479, 73770, 129547, 35390, 7301, + 37318, 107441, 45540, 89189, 116809, 98763, 77827, 82906, 86323, + 5043, 40530, 35486, 24454, 64147, 22581, 127179, 28982, 25813, + 111759, 25875, 67144, 65925, 119945, 65987, 38832, 53596, 102445, + 109559, 69211, 82095, 14947, 117172, 24490, 3989, 530, 105261, + 20337, 23270, 104501, 40575, 111253, 100384, 110528, 20256, 52349, + 44932, 88144, 99082, 32981, 62102, 48256, 120608, 109895, 109835, + 5800, 110786, 116117, 15839, 97802, 39097, 461, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 7995, 85394, 119550, 88325, 48629, 56321, 47267, 70538, 124755, + 114206, 127949, 66234, 97879, 10116, 75346, 116315, 115411, 105906, + 104887, 101066, 124353, 105454, 127841, 24753, 117972, 83013, 23257, + 113424, 62775, 50508, 129544, 56423, 26208, 31226, 41958, 53567, + 44587, 99425, 44966, 34907, 58764, 54870, 85649, 49714, 43361, + 76738, 122439, 125767, 122331, 53573, 6701, 46414, 130886, 89782, + 95205, 67851, 93842, 96807, 57984, 70991, 49875, 67966, 100160, + 53735, 110509, 94681, 29919, 100667, 31561, 122128, 14547, 37662, + 81267, 127530, 117629, 62865, 45977, 80557, 102078, 7436, 34725, + 45975, 64707, 70771, 100145, 81892, 54893, 84306, 1885, 95921, + 84368, 11000, 119123, 35258, 123460, 71345, 38926, 122268, 62525, + 95283, 31822, 110123, 69419, 59439, 126346, 114333, 45092, 55125, + 41949, 99782, 5381, 92466, 79905, 99415, 99944, 3002, 115735, + 82477, 119373, 16392, 56646, 129157, 120981, 70091, 117510, 32163, + 77816, 39738, 79253, 31759, 34433, 38342, 64985, 89927, 47371, + 116134, 97731, 65556, 67396, 115184, 38875, 63741, 95486, 113931, + 113486, 14243, 9757, 94037, 117368, 14744, 97227, 47217, 89316, + 113910, 31862, 84729, 27111, 24920, 4721, 57424, 63294, 105420, + 70146, 45272, 113763, 128548, 130854, 38917, 94062, 130079, 79141, + 24590, 128415, 123042, 16028, 129234, 115403, 111824, 24030, 97272, + 98332, 72842, 105893, 113852, 20216, 73952, 51974, 34840, 63012, + 82468, 85351, 84814, 43057, 87894, 27393, 129833, 69731, 99955, + 34229, 102752, 81848, 58514, 8339, 56606, 33113, 89062, 78126, + 42786, 18705, 14075, 31097, 43319, 127767, 5814, 110148, 27445, + 4752, 28880, 96452, 37793, 26660, 4458, 96852, 84284, 28955, + 86794, 40007, 22025, 48903, 31630, 12571, 71804, 32197, 97091, + 16979, 65070, 112324, 13854, 102150, 46274, 56, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + // ========================================================================= + // ou_dec 正确性测试 + // + // 流程: + // 1. 将 ou_T(base-2^17)转为 base-2^64(OU_T_EXP_LIMBS=4 limb) + // 2. 建立 G/H 预计算表,加密 BATCH 个已知明文 → d_ct(FMLM 域) + // 3. 构造 d_t_exp(batch 个相同的 t,共享私钥) + // 4. GPU:ou_dec(d_ct, d_t_exp) → d_dec(c^t mod n,普通域) + // 5. D2H,写入 ou_dec_test.txt 供 Python 验证 + // + // 注:GPU 输出为 c^t mod n;验证脚本对其 mod p² 即得 c^t mod p², + // 再经 L 函数和 gp_inv 可还原明文 m。 + // ========================================================================= + { + const int BATCH = 200000; // 解密用例数(t 256-bit,单次 FMLE 约数 ms) + + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t m_bytes = (size_t)BATCH * OU_EXP_LIMBS * sizeof(uint64_t); + const size_t rp_bytes = (size_t)BATCH * OU_HR_EXP_LIMBS * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + const size_t texp_bytes = (size_t)BATCH * OU_T_EXP_LIMBS * sizeof(uint64_t); + + // ── Step 1:ou_T(base-2^17,256 limb)→ t_64(base-2^64,4 limb)──────── + // ou_T[i] 代表第 i 个 17-bit 数字,比特偏移 = i*17 + // 映射到 base-2^64 第 j 号 limb(j = bit_offset/64),位移 = bit_offset%64 + uint64_t t_64[OU_T_EXP_LIMBS] = {0}; + for (int i = 0; i < ARR_LEN; i++) { + if (ou_T[i] == 0) continue; + int bit_offset = i * BASE_BITS; + int limb_idx = bit_offset / 64; + int bit_shift = bit_offset % 64; + if (limb_idx < OU_T_EXP_LIMBS) { + t_64[limb_idx] |= (uint64_t)ou_T[i] << bit_shift; + // 若 17 bit 跨越 64-bit 边界,高位部分写入下一个 limb + if (bit_shift + BASE_BITS > 64 && limb_idx + 1 < OU_T_EXP_LIMBS) + t_64[limb_idx + 1] |= (uint64_t)ou_T[i] >> (64 - bit_shift); + } + } + printf("[ou_dec] t_64 = {%llu, %llu, %llu, %llu}\n", + (unsigned long long)t_64[0], (unsigned long long)t_64[1], + (unsigned long long)t_64[2], (unsigned long long)t_64[3]); + + // ── Step 2:建立 G/H 表,加密 BATCH 个已知明文 ────────────────────────── + srand(42u); + + uint64_t *h_G_tbl3 = (uint64_t *)malloc(tbl_bytes); + uint64_t *h_H_tbl3 = (uint64_t *)malloc(tbl_bytes); + generate_G_table(Modn, ou_G, h_G_tbl3); + generate_H_table(Modn, ou_H, h_H_tbl3); + + uint64_t *d_G_tbl3 = nullptr, *d_H_tbl3 = nullptr; + CUDA_CHECK(cudaMalloc(&d_G_tbl3, tbl_bytes)); + CUDA_CHECK(cudaMalloc(&d_H_tbl3, tbl_bytes)); + CUDA_CHECK( + cudaMemcpy(d_G_tbl3, h_G_tbl3, tbl_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_H_tbl3, h_H_tbl3, tbl_bytes, cudaMemcpyHostToDevice)); + free(h_G_tbl3); + free(h_H_tbl3); + + // 生成明文 m(最多 1362 bit) + uint64_t *h_m3 = (uint64_t *)calloc(BATCH * OU_EXP_LIMBS, sizeof(uint64_t)); + for (int b = 0; b < BATCH; b++) { + uint64_t *m = h_m3 + (size_t)b * OU_EXP_LIMBS; + for (int j = 0; j < 21; j++) + m[j] = ((uint64_t)(rand() & 0x7FFF)) | + ((uint64_t)(rand() & 0x7FFF) << 15) | + ((uint64_t)(rand() & 0x7FFF) << 30) | + ((uint64_t)(rand() & 0x7FFF) << 45) | + ((uint64_t)(rand() & 0xF) << 60); + m[21] = (uint64_t)(rand() & 0x3FFFF); + } + + uint64_t *h_rp3 = (uint64_t *)malloc(rp_bytes); + ou_gen_r_prime(ou_p, BATCH, h_rp3); + + uint64_t *d_m3 = nullptr, *d_rp3 = nullptr, *d_ct3 = nullptr; + CUDA_CHECK(cudaMalloc(&d_m3, m_bytes)); + CUDA_CHECK(cudaMalloc(&d_rp3, rp_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct3, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_m3, h_m3, m_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_rp3, h_rp3, rp_bytes, cudaMemcpyHostToDevice)); + + ou_encrypt(d_G_tbl3, d_H_tbl3, d_m3, d_rp3, d_ct3, BATCH, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 3:构造 h_t_exp(每条用例共享同一个 t)───────────────────────── + uint64_t *h_t_exp = (uint64_t *)malloc(texp_bytes); + for (int b = 0; b < BATCH; b++) + for (int j = 0; j < OU_T_EXP_LIMBS; j++) + h_t_exp[(size_t)b * OU_T_EXP_LIMBS + j] = t_64[j]; + + // 先把加密结果从 GPU 取回,用于 H2D 重新上传计时和写文件 + uint64_t *h_ct3 = (uint64_t *)malloc(ct_bytes); + CUDA_CHECK(cudaMemcpy(h_ct3, d_ct3, ct_bytes, cudaMemcpyDeviceToHost)); + + uint64_t *d_ct3_in = nullptr, *d_t_exp = nullptr, *d_dec = nullptr; + CUDA_CHECK(cudaMalloc(&d_ct3_in, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_t_exp, texp_bytes)); + CUDA_CHECK(cudaMalloc(&d_dec, ct_bytes)); + + // ── Step 4:H2D 计时(密文 + 指数 t → GPU)───────────────────────────── + // 模拟真实解密场景:密文和指数均从主机端上传 + cudaEvent_t ev_h2d_s, ev_h2d_e; + CUDA_CHECK(cudaEventCreate(&ev_h2d_s)); + CUDA_CHECK(cudaEventCreate(&ev_h2d_e)); + CUDA_CHECK(cudaEventRecord(ev_h2d_s)); + CUDA_CHECK(cudaMemcpy(d_ct3_in, h_ct3, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_t_exp, h_t_exp, texp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev_h2d_e)); + CUDA_CHECK(cudaEventSynchronize(ev_h2d_e)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev_h2d_s, ev_h2d_e)); + CUDA_CHECK(cudaEventDestroy(ev_h2d_s)); + CUDA_CHECK(cudaEventDestroy(ev_h2d_e)); + + // ── Step 5:GPU ou_dec 计时 ─────────────────────────────────────────── + float ms_dec = ou_dec(d_ct3_in, d_t_exp, d_dec, BATCH, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 6:D2H 计时(解密结果 → 主机)────────────────────────────────── + uint64_t *h_dec = (uint64_t *)malloc(ct_bytes); + cudaEvent_t ev_d2h_s, ev_d2h_e; + CUDA_CHECK(cudaEventCreate(&ev_d2h_s)); + CUDA_CHECK(cudaEventCreate(&ev_d2h_e)); + CUDA_CHECK(cudaEventRecord(ev_d2h_s)); + CUDA_CHECK(cudaMemcpy(h_dec, d_dec, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev_d2h_e)); + CUDA_CHECK(cudaEventSynchronize(ev_d2h_e)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev_d2h_s, ev_d2h_e)); + CUDA_CHECK(cudaEventDestroy(ev_d2h_s)); + CUDA_CHECK(cudaEventDestroy(ev_d2h_e)); + + // ── Step 7:打印计时(单位:us)────────────────────────────────────────── + // H2D 数据量:ct_bytes(密文)+ texp_bytes(指数 t) + // D2H 数据量:ct_bytes(解密结果) + printf("\n========== ou_dec test ==========\n"); + printf(" batch=%d OU_T_TAU=%d ARR_LEN=%d\n", BATCH, OU_T_TAU, ARR_LEN); + printf(" H2D 数据量 : ct=%zu B t=%zu B 合计=%zu B\n", ct_bytes, + texp_bytes, ct_bytes + texp_bytes); + printf(" D2H 数据量 : %zu B\n", ct_bytes); + printf(" ------------------------------------------\n"); + printf(" [H2D ct+t] : %10.2f us (%8.2f us/op)\n", ms_h2d * 1000.f, + ms_h2d * 1000.f / BATCH); + printf(" [GPU ou_dec] : %10.2f us (%8.2f us/op)\n", ms_dec * 1000.f, + ms_dec * 1000.f / BATCH); + printf(" [D2H result] : %10.2f us (%8.2f us/op)\n", ms_d2h * 1000.f, + ms_d2h * 1000.f / BATCH); + printf(" ------------------------------------------\n"); + printf(" [Total] : %10.2f us (%8.2f us/op)\n", + (ms_h2d + ms_dec + ms_d2h) * 1000.f, + (ms_h2d + ms_dec + ms_d2h) * 1000.f / BATCH); + printf("=================================\n\n"); + /* + // 格式 ou_dec_test.txt: + // 第 1 行(注释):元信息 + // 第 2 行 N: OU 模数 n(256 limb,base-2^17) + // 第 3 行 p: 素数 p (256 limb,base-2^17) + // 第 4 行 T64: 解密指数 t(OU_T_EXP_LIMBS=4 limb,base-2^64) + // 第 5 行 GP: gp_inv (256 limb,base-2^17) + // 每用例 3 行: + // M: 明文 m(OU_EXP_LIMBS=22 limb,base-2^64) + // CT: 密文 c̃(256 limb,base-2^17,FMLM 域) + // DEC: c^t mod n(256 limb,base-2^17,普通域) + FILE *fp = fopen("ou_dec_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_dec_test.txt\n"); + } else { + fprintf(fp, + "# ou_dec test BATCH=%d ARR_LEN=%d BASE_BITS=%d" + " OU_T_TAU=%d OU_T_EXP_LIMBS=%d\n", + BATCH, ARR_LEN, BASE_BITS, OU_T_TAU, OU_T_EXP_LIMBS); + // N: + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + // p: + fprintf(fp, "p:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_p[j]); + fprintf(fp, "\n"); + // T64: 解密指数 t(base-2^64) + fprintf(fp, "T64:"); + for (int j = 0; j < OU_T_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)t_64[j]); + fprintf(fp, "\n"); + // GP: gp_inv(base-2^17) + fprintf(fp, "GP:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_gp_inv[j]); + fprintf(fp, "\n"); + // 逐用例写 M / CT / DEC + for (int b = 0; b < BATCH; b++) { + const uint64_t *m = h_m3 + (size_t)b * OU_EXP_LIMBS; + const uint64_t *ct = h_ct3 + (size_t)b * ARR_LEN; + const uint64_t *dec = h_dec + (size_t)b * ARR_LEN; + fprintf(fp, "M:"); + for (int j = 0; j < OU_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)m[j]); + fprintf(fp, "\n"); + fprintf(fp, "CT:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ct[j]); + fprintf(fp, "\n"); + fprintf(fp, "DEC:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)dec[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_dec test] 测试数据已写入 ou_dec_test.txt\n"); + } + */ + // ── 释放 ────────────────────────────────────────────────────────────── + free(h_m3); + free(h_rp3); + free(h_ct3); + free(h_dec); + free(h_t_exp); + cudaFree(d_m3); + cudaFree(d_rp3); + cudaFree(d_ct3); + cudaFree(d_G_tbl3); + cudaFree(d_H_tbl3); + cudaFree(d_ct3_in); + cudaFree(d_t_exp); + cudaFree(d_dec); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_dec_complete.cu b/heu/library/algorithms/ou_new/ou_dec_complete.cu new file mode 100644 index 0000000..e23b665 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_dec_complete.cu @@ -0,0 +1,19980 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12420598007131, 404862574961, 1489476715458592911, + 16957267883830147813, 12043402377674848005, 16493683929612125450, + 6784216004999364667, 1572187082686061104, 12340520880335831418, + 9331149018221532047, 14282361491840612423, 8815005708360508971, + 3015263268809804842, 17865818108905781022, 9920159502644748834, + 16663441635849680526, 16754665591114057436, 18269507414440088358, + 10584069732924195264, 12369982964838682500, 17245146333661260496, + 14241460736918561983, 1495228790312205286, 650596912814140113, + 13561313259753582639, 12825917001300214751, 17648203335955277871, + 4937017841769177676, 8762411814459186696, 10916964928297372446, + 7934650143980081733, 16270299917683438672, 10599916200641161267, + 8153067147237791839, 3263540562937225549, 12427502011302527075, + 233468242765771385, 16905948202777045661, 7718378239207545448, + 12556027611679852532, 14355997036124831232, 10480393087867199263, + 8928529658741694940, 5329501958372121449, 14697840141414380795, + 4109255465868101899, 2189093870568573914, 2015298188623315682, + 15846288464350737625, 6154716281955634074, 14926805960051698119, + 6358221390244429391, 565479903706287841, 544001290941279643, + 8964697110262868308, 18143181760096837871, 4558614659708563008, + 13794757981490580748, 8464822194361423649, 3215767361281822704, + 3162014795756580172, 4046965963624353869, 13045127494050167872, + 2499196209129003472, 9910258488503570389, 17361620693833999278, + 9909125957613821241, 9160076386865281785, 1124681048989805109, + 13488748519780692092, 10062780796386632218, 1229663950817871070, + 9195972272372452954, 16317888560940382210, 7613588099474937282, + 4817436175622690809, 2044827062874516597, 10883598385799752225, + 15983746154563805830, 15305213912578498862, 15429027565390837948, + 15165749666894862326, 17094752132458668719, 17119153326423120334, + 16418978274070441262, 4020071962619695082, 1332089103620365082, + 2185942973674898504, 10623386207974578633, 17221670917682379713, + 1728235680214257322, 4351774997312820994, 13118214034020006923, + 17372706552005974065, 3635082452886395532, 9785476544995844000, + 5239966118419100751, 11446467251832755986, 8441993571518844572, + 17814051721047964692, 6635316643940291166, 591992533123407590, + 5392319256964754787, 6876233685503189066, 12560692609570806257, + 15922352178387947359, 10408865844002002467, 8334121862371121723, + 12325068481620777235, 16272399208182333772, 95913092914333903, + 9337511878088658138, 3620707693508746669, 2410461826666169902, + 15825609924068456517, 14631303029849624577, 12745725307310343054, + 11589389212418519430, 2402535340645687114, 10460197248544579201, + 2150167009874019400, 6579654444057751335, 10629733729325322988, + 15336435747150065712, 2033376745185295929, 5536738554015879544, + 6536459867767524988, 17547258140832520218, 16277632797518007476, + 17414093779241440205, 15172877332690969372, 11510709146567555164, + 4030987138388375538, 4545399202703904429, 3071928236698087396, + 6235823038589723214, 1385517152964008659, 5315971425862540039, + 737213825926788813, 12647942483396120777, 11729199013777796647, + 16701826799686034287, 10264009214043680734, 15712857482319384195, + 9683635541859048778, 12510245784932525987, 11453973523166319383, + 17641407202450751284, 7076031528159713278, 4418210339976245128, + 7042975543133840791, 2687978692681798677, 6620884024635649052, + 18356048427182277532, 12306609112475240772, 9532567709512160247, + 18330814574547253380, 1570356219924936203, 8589904627747755897, + 15801551039736301865, 14430949397822386754, 17457931947828226443, + 17066621315075642576, 5110941585241228080, 3113435720621532772, + 9109366487892874842, 6870609824263132877, 12909242032068451345, + 9978696359698762281, 2334245554206682016, 5474675901673095882, + 6791983755440567441, 7940389551176083495, 8857160764169799470, + 5079795210984134207, 13270967636848123637, 2817746455641616989, + 6823734589021892698, 6695712168139536841, 10011091548170035449, + 1880778261691426396, 4901098797777619034, 10275511334773938844, + 14807980597086261468, 2140936956496274302, 16031412187684387227, + 13312372442173260977, 3698443504872582944, 13391615949904799953, + 9911497060502159398, 6476394081735368123, 10884190273162182630, + 16812094898227487679, 17541511080948430118, 608781942072315334, + 11293244422430215336, 11785770688631267218, 1050397619038471432, + 8886815311739454498, 4768360092769081595, 6853419982646256117, + 4533330592021381799, 4065300292447750493, 5995762058657812513, + 5850499564985481696, 17494639705373259696, 5311542451880315192, + 10494352427384472421, 7053555067647955415, 13249062068286530876, + 4950095133067012336, 978416016041957203, 16132684490165566392, + 9461678038864020399, 13953140033946412303, 3128510702558997798, + 9989688919461798931, 7729804782429579018, 14657677073053593359, + 2033538150495053663, 18287519723785336581, 15945421906051418344, + 827199110667409676, 2705644704721997071, 7966095645914939971, + 16210256519322987223, 13500279664518420546, 6403250311096390047, + 15684661198519501341, 17612977561176850685, 1595673854820173645, + 11355330109781124768, 15152054633977714227, 7594126125750130131, + 1770764465613758168, 16103674937332423683, 12989244865424659639, + 985924609341537427, 17869081554936260320, 13756399484785420256, + 9417665840560873340, 6596619827292141536, 9282779747668579711, + 14087089070266996746, 16797541746528108966, 3649008588426316857, + 305957104569710731, 16307164143397236406, 8758406156783513520, + 8985939822111752351, 9823364479239409710, 14242788910385080985, + 16257804781363675711, 3460875886956782926, 5792587040463753957, + 8058808855494344097}; +const uint64_t con_modn_shoup[256] = { + 17258949916450502562, 8888112940397698219, 713580727021815326, + 7267352695482246207, 16882106641962718850, 4484957681626972705, + 11436540044429239973, 11422524016867871168, 14020805313791955351, + 3223598656225742213, 5973835375502394258, 18389866180389685123, + 16986950504460752843, 14151661301513655720, 17077263121249724353, + 9230648276690066133, 14180947083865490845, 9351212926919619729, + 14132401609129010813, 5094848903632022188, 9701476928191356694, + 2266122304659313991, 9642600553848677189, 17694947392610440444, + 4166068824698342286, 5461002667202655000, 12457224361386750838, + 17183384576469292838, 4577122507392282256, 212990517287292298, + 9593917155195774666, 16026775048324125596, 16942601617377608128, + 14517179166542822181, 5974534123485502056, 3098739027972382032, + 10256410151333955680, 17520563159797137486, 1437139722912674326, + 8992878724252091890, 12528748759234609070, 8869909966108752023, + 921545360850748754, 9072672431128383783, 5738738820990719155, + 14467747599114998535, 17288556119140904797, 10287232211619750677, + 14260418889536810543, 14236170919296794224, 3150044698371016975, + 3734472856964467879, 17367148054045259108, 8704684147659221983, + 2508292278448591736, 8252783960354095393, 5505177864158443900, + 6073691067055527962, 1797426869555915200, 10910566681039151494, + 7942925665011098265, 2975836443917988571, 12168486674970259490, + 16110198076303607574, 3176765343615554343, 16298611667673954183, + 17549972852929401068, 11301222901206006551, 16909965679872285631, + 4718899407226931290, 3249759906455460509, 12902044256019035476, + 2933343404927512513, 12635530155166610452, 1019795786854894305, + 5558319569785986415, 1736084326860448187, 1046201637077468719, + 8212969089240241408, 14337127378294037853, 386901202666220852, + 7007901116396992081, 12993354667310087947, 10362596985035361337, + 2632315220458003599, 15774429557071893885, 16402946062094934668, + 1455058537672538730, 10934206267514764178, 8653495955263034803, + 15626427345777561198, 9949509055664635150, 13551873570797002283, + 8464454269390886841, 17029555676738610993, 2576180826541960979, + 9277239235861014619, 5949389895883929521, 5961129371502370384, + 205015922583922578, 10998263302456627347, 630682958888475996, + 8272932531190805541, 11739326145926911850, 3711218700419121610, + 1995808726167056948, 9541493218673654357, 10947701895562230579, + 13307121742110758926, 5577478578055132785, 17291845771099354096, + 18386024476421024505, 12682871500298847381, 11130873186526484524, + 9511555001775281238, 12573445966195088218, 8735156148113491387, + 11986977738879871749, 3957140364191407682, 9447015894594054903, + 15075232292569908924, 3903378683888054983, 18108762427411217092, + 3848515571709448978, 9812945716174782266, 1439686790983251294, + 16833581092072042070, 16210195608156458739, 5586553670046771142, + 16285469091149528676, 9398773791615283024, 6409294867327805667, + 7451696515003176156, 9352800728836756410, 13006499844621233701, + 8943861815085160245, 9678575573847830426, 8255945162964001220, + 10730739327269562739, 5334149872806451563, 17598838495642584396, + 16369098332122517680, 17955871179034626630, 14641785369093985684, + 14862588269370659671, 18285743176298903762, 13037469746309732440, + 14680321993850845543, 4817891414280274891, 5149080315255804853, + 2128397328110922197, 10589038400063582404, 7361242319334319176, + 2633416174015188179, 8731722470221646364, 17695241939569347999, + 511281303145693578, 3915229765025655898, 12611201828639463913, + 12923773568272819998, 6599022632397876071, 13420679683499938789, + 2049055346114188916, 12969132418156712259, 4309319847942653387, + 18244990947493391329, 4321703802367122695, 6192941661180839271, + 3809220907288466934, 9686397595036380761, 5512567166054515334, + 7602079309929874744, 13912271487139773395, 14184657929002112052, + 14576671517535568717, 10645713420980679390, 6189820458973792135, + 11814662608890796689, 14130844520328466724, 5920079432947567926, + 11563681922574043697, 1737013972124835278, 9154718647461864771, + 4872088247015674989, 3298657241789078374, 4274525812376499266, + 2844739318930852390, 5423787114207982915, 711603148563500498, + 5237712164936082793, 16538115545248024540, 6726637234751909654, + 12914942507138411381, 13794083152765605836, 8092188286693750926, + 3585046709821797374, 9022641750310238763, 465420587122630117, + 13564884715383175347, 11129511873502417680, 5164151854448859979, + 15029903335193380621, 10868746496347154392, 4435651212450171864, + 13217684022763318278, 5933919845722166548, 3356965223531861545, + 5310913073599813054, 15054780849600338193, 6396900692144724485, + 7035542304377335696, 10193072637584812570, 13651162572970216008, + 3375517112149222934, 4439239668198625265, 4272637244444960409, + 12817026244019634782, 69354263255062928, 15302789662880427742, + 347862105615759955, 17769797205451930634, 9642432860176744960, + 12464642972155630307, 15664388251667466866, 11631517894848881873, + 17036907573636512589, 7095039048837509769, 17513417735648162476, + 17481830685724223605, 17280286146669482802, 2842388552256186635, + 3920667164872509047, 12274188700851890415, 10745285319594444214, + 1477783310694359069, 9823805546543362486, 13036068216114647306, + 2630587855422542122, 11133180172685530573, 5498501119950892568, + 12148543347986984583, 8699653759414714944, 12680167660387601338, + 7199244163763404941, 11286265673566327732, 16227882229914399460, + 13238315363757133297, 15029696251503760467, 3493558449506003864, + 1655355875307187670, 2509903151438279534, 16562697969364590150, + 16509371521602835055, 6068654425169380909, 12704430300130691212, + 1645898033077800833}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define OU_T_TAU 256 // 解密指数 t 的 bit 数(p-1 的大素因子,~256 bit) +#define OU_T_EXP_LIMBS (OU_T_TAU / 64) // = 4,每个 t 占 4 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +static void ou_gen_r_prime( + const uint64_t * /*ou_p_limbs17*/, // 保留参数兼容调用处,128-bit + // 时不需要 p + int batch, uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================================= +// ou_encrypt +// +// 功能:批量 OU 加密(正明文) +// c̃[i] = FMLM(G^{m_i}, H^{r'_i}) mod n (FMLM 域输出) +// +// 数学依据: +// OU 加密公式:c = G^m · H^r mod n +// FMLM 域语义:FMLM(Ã, B̃) = ÷B̃·R⁻¹ mod n +// 因为 Getgp/Getrn 输出 FMLM 域(乘了 R),所以: +// FMLM(G^m·R, H^r'·R) = G^m·H^r'·R mod n ✓ +// +// 三步流程: +// Step① Getgp :G 预计算表 × m_i → G^{m_i} (FMLM 域)→ d_ct_out +// Step② Getrn :H 预计算表 × r'_i → H^{r'_i}(FMLM 域)→ d_Hr(临时) +// Step③ XYfixWarpVector:d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) +// +// 调用前准备: +// generate_G_table(Modn, ou_G, h_G_table) 并上传 → d_G_table +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM +// 域,generate_G_table 输出) d_H_table [TABLE_SIZE × ARR_LEN] H +// 预计算表(FMLM 域,generate_H_table 输出) d_m_batch [batch × +// OU_EXP_LIMBS] 明文 m(base-2^64 小端序,OU_EXP_LIMBS=22) +// d_r_prime_batch [batch × OU_HR_EXP_LIMBS] 随机指数 r'(OU_HR_TAU=128 +// bit,OU_HR_EXP_LIMBS=2) d_ct_out [batch × ARR_LEN] 输出密文(FMLM +// 域) batch 明文数量 +// +// 返回:GPU 端 Step①②③ 总耗时(ms) +// ============================================================================= +float ou_encrypt(const uint64_t *d_G_table, const uint64_t *d_H_table, + const uint64_t *d_m_batch, // [batch × OU_EXP_LIMBS] + const uint64_t *d_r_prime_batch, // [batch × OU_HR_EXP_LIMBS] + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM 域) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getgp — G^{m_i} mod n(FMLM 域)→ d_ct_out + Getgp<<>>( + d_G_table, d_m_batch, OU_TAU, d_ct_out, batch, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step③:XYfixWarpVector — d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) + // 结果:c̃[i] = G^{m_i} · H^{r'_i} · R mod n(FMLM 域密文) + XYfixWarpVector<<>>( + d_ct_out, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================= +// ou_dec +// +// 功能:GPU 批量 OU 解密第一步——模幂 c^t mod n(普通域输出) +// +// OU 完整解密流程: +// Step①(本函数): c^t mod n ← FMLE_mod3_Kernel,NTT 参数为 n 模数 +// Step②(CPU 端): (result) mod p² ← 因 n=p²·q,c^t mod n 再 mod p² = c^t +// mod p² Step③(CPU 端): m = L(c') · gp_inv mod p,其中 L(x) = (x-1)/p +// +// 输入 d_c_tilde 为 FMLM 域密文 c̃ = c·R mod n; +// FMLE_mod3_Kernel 跳过 Step8(输入已是 FMLM 域), +// 保留 Step13(乘 R⁻¹,还原为普通域),输出 c^t mod n。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(FMLM 域,mod n) +// d_t_exp [batch × OU_T_EXP_LIMBS] 指数 t(压缩格式:OU_T_EXP_LIMBS=4 个 +// uint64_t, +// 小端序,位 i 在 limb[i/64] 的第 i%64 +// 位) +// d_output [batch × ARR_LEN] 输出:c^t mod n(普通域,base-2^17) +// batch 批大小 +// d_r0 [ARR_LEN] r₀ = (2^(ARR_LEN×BASE_BITS)-1) mod n +// (FMLM 单位元,即蒙哥马利域中的 1) +// 其余为 NTT 参数(n 模数,与 ou_encrypt 完全一致) +// +// 返回:GPU 端到端耗时(ms) +// ============================================================= +float ou_dec(const uint64_t *d_c_tilde, const uint64_t *d_t_exp, + uint64_t *d_output, int batch, const uint64_t *d_r0, + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * OU_T_EXP_LIMBS; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_t_exp + exp_off, OU_T_TAU, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================================ +// ou_broadcast_kernel: 将单份 src[ARR_LEN] 广播到 dst[batch × ARR_LEN] +// ============================================================================ +__global__ void ou_broadcast_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================================ +// ou_compute_p2_kernel: 单 instance CGBN 计算 p² = p * p +// 启动参数: <<<1, INV_TPI>>> +// ============================================================================ +__global__ void ou_compute_p2_kernel(cgbn_error_report_t *report, + inv_bn_mem_t *d_p2, inv_bn_mem_t *d_p) { + int instance_id = (blockIdx.x * blockDim.x + threadIdx.x) / INV_TPI; + if (instance_id != 0) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t p, p2; + cgbn_load(env, p, d_p); + cgbn_mul(env, p2, p, p); // p ≈ 1364 bit, p² ≈ 2728 bit < 4096 bit,不截断 + cgbn_store(env, d_p2, p2); +} + +// ============================================================================ +// ou_L_kernel: GPU CGBN 批量计算 L = (c^t mod p² − 1) / p +// 每 INV_TPI(=32) 线程处理一个实例 +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_L_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_L, // [batch] 输出 + inv_bn_mem_t *d_ct, // [batch] 输入:c^t(CGBN 格式) + inv_bn_mem_t *d_p2, // [1] 输入:p²(常量,所有实例共用) + inv_bn_mem_t *d_p, // [1] 输入:p (常量,所有实例共用) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t ct, p2, p_bn, tmp, L; + + cgbn_load(env, ct, d_ct + instance_id); + cgbn_load(env, p2, d_p2); // 所有实例共用同一个 p² + cgbn_load(env, p_bn, d_p); + + cgbn_rem(env, tmp, ct, p2); // tmp = c^t mod p² + cgbn_sub_ui32(env, tmp, tmp, + 1); // tmp = tmp − 1 (OU 保证 c^t ≡ 1 mod p,故 tmp ≥ 1) + cgbn_div(env, L, tmp, p_bn); // L = tmp / p (精确整除) + + cgbn_store(env, d_L + instance_id, L); +} + +// ============================================================================ +// ou_modp_kernel: GPU CGBN 批量计算 m = prod mod p +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_modp_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_m, // [batch] 输出 + inv_bn_mem_t *d_prod, // [batch] 输入:L × gp_inv mod n + inv_bn_mem_t *d_p, // [1] 输入:p(常量) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t prod, p_bn, m; + + cgbn_load(env, prod, d_prod + instance_id); + cgbn_load(env, p_bn, d_p); + cgbn_rem(env, m, prod, p_bn); + + cgbn_store(env, d_m + instance_id, m); +} + +// ============================================================================ +// ou_dec_complete: 完整 OU 解密 +// +// Step① ou_dec : d_c_tilde → d_ct_plain (c^t mod n,普通域) +// Step② format_to_cgbn_kernel : d_ct_plain → d_ct_cgbn +// Step③ ou_L_kernel : d_ct_cgbn → d_L_cgbn (L=(c^t mod +// p²−1)/p) Step④ format_from_cgbn_kernel : d_L_cgbn → d_L_b17 Step⑤ +// XYfixWarpROneVector×1 : d_gp_inv → 蒙哥马利域(原地,单份) Step⑥ +// ou_broadcast_kernel : d_gp_inv → d_gp_inv_batch(batch 份) Step⑦ +// XYfixWarpVector : d_L_b17 × d_gp_inv_batch → L×gp_inv mod n Step⑧ +// format_to_cgbn_kernel : d_L_b17 → d_prod_cgbn Step⑨ ou_modp_kernel : +// d_prod_cgbn → d_m_cgbn (mod p) Step⑩ format_from_cgbn_kernel : d_m_cgbn +// → d_m_out +// +// 返回: GPU 全流程耗时(ms,含 ou_dec 内部时间) +// ============================================================================ +float ou_dec_complete( + const uint64_t *d_c_tilde, // [batch × ARR_LEN] FMLM 域密文 + const uint64_t *d_t_exp, // [batch × OU_T_EXP_LIMBS] 指数 t + uint64_t *d_m_out, // [batch × ARR_LEN] 输出:明文 m(base-2^17) + int batch, + const uint64_t *h_p, // p(主机端,base-2^17) + const uint64_t *h_gp_inv, // gp_inv(主机端,base-2^17) + const uint64_t *d_r0, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + uint64_t *d_ctR, // CT(R²),传给 XYfixWarpROneVector + uint64_t *d_nctR // NCT(R²),传给 XYfixWarpROneVector +) { + const size_t ct_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + const size_t cgbn_bytes = (size_t)batch * sizeof(inv_bn_mem_t); + const size_t one_cgbn = sizeof(inv_bn_mem_t); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + // ── 申请设备临时内存 ───────────────────────────────────────────────────── + uint64_t *d_ct_plain = nullptr; + inv_bn_mem_t *d_ct_cgbn = nullptr; + inv_bn_mem_t *d_p_cgbn = nullptr; + inv_bn_mem_t *d_p2_cgbn = nullptr; + inv_bn_mem_t *d_L_cgbn = nullptr; + uint64_t *d_L_b17 = nullptr; + uint64_t *d_gp_inv = nullptr; + uint64_t *d_gp_inv_batch = nullptr; + inv_bn_mem_t *d_prod_cgbn = nullptr; + inv_bn_mem_t *d_m_cgbn = nullptr; + + CUDA_CHECK(cudaMalloc(&d_ct_plain, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_p_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_p2_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_L_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_b17, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_gp_inv, ARR_LEN * sizeof(uint64_t))); + CUDA_CHECK(cudaMalloc(&d_gp_inv_batch, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_prod_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_cgbn, cgbn_bytes)); + + // ── Step①: c^t mod n ──────────────────────────────────────────────────── + ou_dec(d_c_tilde, d_t_exp, d_ct_plain, batch, d_r0, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + // ou_dec 内部使用多流,需同步后才能进行后续 CGBN 操作 + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 上传 p、gp_inv;GPU 计算 p² ────────────────────────────────────────── + { + inv_bn_mem_t h_p_cgbn; + bn17_to_cgbn_mem(h_p, &h_p_cgbn); + CUDA_CHECK( + cudaMemcpy(d_p_cgbn, &h_p_cgbn, one_cgbn, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_gp_inv, h_gp_inv, ARR_LEN * sizeof(uint64_t), + cudaMemcpyHostToDevice)); + + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + ou_compute_p2_kernel<<<1, INV_TPI>>>(report, d_p2_cgbn, d_p_cgbn); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] p² 计算 CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step②: base-2^17 → CGBN ───────────────────────────────────────────── + format_to_cgbn_kernel<<>>(d_ct_cgbn, d_ct_plain, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③: L = (c^t mod p² − 1) / p ──────────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_L_kernel<<>>(report, d_L_cgbn, d_ct_cgbn, d_p2_cgbn, + d_p_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_L_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step④: CGBN → base-2^17 (L) ──────────────────────────────────────── + format_from_cgbn_kernel<<>>(d_L_b17, d_L_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑤: gp_inv 转蒙哥马利域(单份)────────────────────────────────── + // FMLM(gp_inv, R²) = gp_inv * R mod n → 蒙哥马利域,结果原地写回 d_gp_inv + XYfixWarpROneVector<<<1, 32>>>( + d_gp_inv, d_ctR, d_nctR, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑥: 广播 gp_inv_mont 到 batch 份 ──────────────────────────────── + { + const int total = batch * ARR_LEN; + const int bthreads = 256; + const int bblocks = (total + bthreads - 1) / bthreads; + ou_broadcast_kernel<<>>(d_gp_inv_batch, d_gp_inv, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step⑦: FMLM(L_plain, gp_inv_mont) = L × gp_inv mod n(普通域)────── + // inout=d_L_b17(普通域),inoutA=d_gp_inv_batch(蒙哥马利域) + // 结果 = L * gp_inv_mont * R^{-1} = L * gp_inv mod n,写回 d_L_b17 + XYfixWarpVector<<>>( + d_L_b17, d_gp_inv_batch, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑧: base-2^17 → CGBN (L × gp_inv mod n) ───────────────────────── + format_to_cgbn_kernel<<>>(d_prod_cgbn, d_L_b17, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑨: m = (L × gp_inv mod n) mod p ──────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_modp_kernel<<>>(report, d_m_cgbn, d_prod_cgbn, d_p_cgbn, + batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_modp_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step⑩: CGBN → base-2^17 (m,写入 d_m_out) ────────────────────────── + format_from_cgbn_kernel<<>>(d_m_out, d_m_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + // ── 释放临时设备内存 ───────────────────────────────────────────────────── + cudaFree(d_ct_plain); + cudaFree(d_ct_cgbn); + cudaFree(d_p_cgbn); + cudaFree(d_p2_cgbn); + cudaFree(d_L_cgbn); + cudaFree(d_L_b17); + cudaFree(d_gp_inv); + cudaFree(d_gp_inv_batch); + cudaFree(d_prod_cgbn); + cudaFree(d_m_cgbn); + + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +int main() { + uint64_t Modn[256] = { + 90459, 46625, 4442, 112182, 114968, 71534, 86507, 127438, 1685, + 52931, 104150, 120548, 110008, 7428, 47921, 17877, 51604, 95338, + 118828, 83898, 128532, 25064, 75321, 17612, 57426, 74847, 91485, + 17341, 32555, 39547, 13559, 24145, 41272, 116851, 19414, 130804, + 122030, 33140, 103122, 68029, 121047, 93139, 39899, 122526, 23084, + 120314, 105769, 34702, 77999, 9851, 111681, 49144, 6962, 87978, + 73280, 9230, 120490, 99393, 36052, 66751, 52136, 26210, 79015, + 53220, 88620, 106517, 78130, 39022, 74644, 83791, 78431, 76879, + 126508, 112379, 112791, 45016, 126862, 98555, 67445, 123550, 78839, + 94946, 54918, 116122, 88926, 128743, 56161, 106652, 89566, 6771, + 65824, 105202, 2315, 108169, 98258, 27489, 25367, 128358, 45812, + 97612, 65111, 13602, 17793, 42341, 58555, 122795, 20331, 7114, + 60301, 88265, 43599, 10109, 18770, 82809, 13834, 117670, 103794, + 49075, 68168, 87021, 105358, 120278, 55703, 7511, 51479, 98172, + 58294, 99843, 55022, 33556, 19683, 122106, 75656, 112874, 46791, + 112145, 106681, 122130, 87140, 32269, 9030, 59222, 342, 75944, + 94103, 130245, 53043, 71941, 64227, 216, 36445, 98379, 21169, + 62160, 91119, 86471, 3598, 23603, 98537, 23157, 124090, 112187, + 67586, 111947, 77293, 12685, 36844, 113775, 50048, 260, 118397, + 88108, 66275, 63944, 42517, 39246, 24220, 115910, 91147, 23953, + 45252, 7345, 75180, 16316, 123347, 38049, 74728, 78637, 21442, + 41317, 90948, 49684, 37151, 40998, 36225, 79231, 81693, 87993, + 55886, 2798, 82414, 110989, 110453, 57672, 14220, 82985, 1843, + 128765, 104271, 125420, 13546, 63807, 90504, 103167, 51426, 33297, + 102375, 75357, 118381, 38152, 11460, 114990, 16009, 63757, 78839, + 46610, 83493, 100684, 103041, 23340, 76918, 82210, 94466, 46728, + 76817, 44869, 84237, 127471, 98069, 60741, 890, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 89416, 26754, 112835, 104869, 63192, 109205, 3104, 29726, 113565, + 56327, 49094, 67751, 77782, 27488, 119631, 71590, 123612, 97756, + 84117, 11401, 55613, 66158, 70476, 101675, 91512, 52471, 103039, + 24137, 47894, 66258, 102515, 73290, 122151, 114889, 128986, 74711, + 74396, 24226, 116902, 16522, 70549, 97827, 20059, 59149, 52686, + 26411, 100945, 78310, 25925, 67014, 79121, 123619, 40826, 66607, + 30092, 89869, 105580, 20910, 93204, 44105, 73778, 21121, 113212, + 113354, 12462, 2188, 4706, 7519, 55812, 128963, 35815, 6339, + 118625, 95153, 80551, 71853, 110048, 116886, 28114, 36427, 99347, + 11391, 127843, 41917, 14312, 58727, 41417, 4720, 129585, 25994, + 106019, 13449, 40561, 107035, 2390, 33535, 7847, 22268, 12977, + 53724, 114048, 27127, 106539, 77244, 114474, 118543, 104524, 97508, + 116569, 73049, 95427, 14163, 96131, 121201, 90712, 20571, 129841, + 128480, 52654, 46075, 61521, 53523, 19041, 42853, 127248, 11120, + 123997, 130413, 56569, 98615, 56998, 99234, 71154, 41850, 65057, + 46995, 104286, 37295, 68580, 38759, 28487, 15348, 68337, 102177, + 78016, 31678, 52663, 1195, 70670, 125162, 16806, 115057, 38759, + 1825, 67883, 105361, 112649, 60917, 88939, 64087, 91926, 11035, + 53857, 46013, 61091, 124646, 128853, 106128, 66902, 108804, 8022, + 25699, 7098, 119054, 99103, 128189, 128398, 1319, 96584, 80228, + 81809, 54282, 6764, 98910, 13622, 81427, 9254, 125751, 67112, + 129121, 60904, 13975, 42521, 71104, 19357, 130027, 15821, 92416, + 81803, 63240, 43145, 90507, 34211, 103714, 62407, 43497, 11360, + 9542, 117006, 82808, 14980, 96632, 5649, 54610, 109476, 15174, + 113385, 48087, 34882, 126953, 67608, 7544, 126322, 75249, 4473, + 4549, 79315, 60258, 82932, 85897, 13892, 119911, 78557, 112723, + 15372, 51919, 45301, 73004, 122707, 121269, 108268, 19504, 49759, + 64648, 50866, 115280, 49268, 9806, 71075, 96540, 114884, 100204, + 79963, 109285, 92516, 38345}; + uint64_t R1[256] = { + 63840, 77643, 38589, 103335, 39211, 83771, 23554, 103602, 13011, + 33265, 99850, 107872, 65932, 9450, 98916, 44944, 70258, 117509, + 29436, 61780, 95269, 57764, 113636, 76719, 51738, 24730, 60149, + 46944, 29753, 55286, 121585, 3994, 70184, 56425, 18607, 115724, + 66839, 44507, 95968, 14278, 79412, 78754, 21044, 58743, 116286, + 99739, 26812, 85232, 130351, 42770, 22378, 65680, 11282, 130626, + 67800, 75341, 43729, 3643, 41124, 90487, 21576, 61026, 35178, + 84381, 65484, 94672, 34715, 116332, 18621, 103526, 117440, 57014, + 7488, 192, 2821, 11532, 82352, 92862, 129569, 103329, 115342, + 91522, 84549, 112318, 128186, 46600, 126196, 12657, 116, 12915, + 15556, 72466, 25692, 116854, 113547, 7800, 66899, 57704, 21718, + 49267, 129570, 121088, 11232, 26664, 76057, 98161, 88201, 110810, + 67711, 57350, 129067, 55009, 66719, 82831, 16394, 45675, 54590, + 128704, 53971, 119083, 90622, 38432, 89295, 6632, 106784, 29141, + 112848, 12094, 121287, 45606, 18268, 123439, 23933, 114383, 91280, + 32929, 75475, 33289, 81297, 101051, 117022, 93617, 78312, 23394, + 80435, 107542, 74032, 67728, 108966, 130175, 112943, 113217, 3110, + 3831, 11764, 90166, 19552, 130175, 21228, 105234, 12469, 127827, + 48250, 34863, 105366, 25850, 91309, 59492, 11089, 37416, 86082, + 43171, 97876, 17564, 124193, 22479, 73770, 129547, 35390, 7301, + 37318, 107441, 45540, 89189, 116809, 98763, 77827, 82906, 86323, + 5043, 40530, 35486, 24454, 64147, 22581, 127179, 28982, 25813, + 111759, 25875, 67144, 65925, 119945, 65987, 38832, 53596, 102445, + 109559, 69211, 82095, 14947, 117172, 24490, 3989, 530, 105261, + 20337, 23270, 104501, 40575, 111253, 100384, 110528, 20256, 52349, + 44932, 88144, 99082, 32981, 62102, 48256, 120608, 109895, 109835, + 5800, 110786, 116117, 15839, 97802, 39097, 461, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 7995, 85394, 119550, 88325, 48629, 56321, 47267, 70538, 124755, + 114206, 127949, 66234, 97879, 10116, 75346, 116315, 115411, 105906, + 104887, 101066, 124353, 105454, 127841, 24753, 117972, 83013, 23257, + 113424, 62775, 50508, 129544, 56423, 26208, 31226, 41958, 53567, + 44587, 99425, 44966, 34907, 58764, 54870, 85649, 49714, 43361, + 76738, 122439, 125767, 122331, 53573, 6701, 46414, 130886, 89782, + 95205, 67851, 93842, 96807, 57984, 70991, 49875, 67966, 100160, + 53735, 110509, 94681, 29919, 100667, 31561, 122128, 14547, 37662, + 81267, 127530, 117629, 62865, 45977, 80557, 102078, 7436, 34725, + 45975, 64707, 70771, 100145, 81892, 54893, 84306, 1885, 95921, + 84368, 11000, 119123, 35258, 123460, 71345, 38926, 122268, 62525, + 95283, 31822, 110123, 69419, 59439, 126346, 114333, 45092, 55125, + 41949, 99782, 5381, 92466, 79905, 99415, 99944, 3002, 115735, + 82477, 119373, 16392, 56646, 129157, 120981, 70091, 117510, 32163, + 77816, 39738, 79253, 31759, 34433, 38342, 64985, 89927, 47371, + 116134, 97731, 65556, 67396, 115184, 38875, 63741, 95486, 113931, + 113486, 14243, 9757, 94037, 117368, 14744, 97227, 47217, 89316, + 113910, 31862, 84729, 27111, 24920, 4721, 57424, 63294, 105420, + 70146, 45272, 113763, 128548, 130854, 38917, 94062, 130079, 79141, + 24590, 128415, 123042, 16028, 129234, 115403, 111824, 24030, 97272, + 98332, 72842, 105893, 113852, 20216, 73952, 51974, 34840, 63012, + 82468, 85351, 84814, 43057, 87894, 27393, 129833, 69731, 99955, + 34229, 102752, 81848, 58514, 8339, 56606, 33113, 89062, 78126, + 42786, 18705, 14075, 31097, 43319, 127767, 5814, 110148, 27445, + 4752, 28880, 96452, 37793, 26660, 4458, 96852, 84284, 28955, + 86794, 40007, 22025, 48903, 31630, 12571, 71804, 32197, 97091, + 16979, 65070, 112324, 13854, 102150, 46274, 56, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + // ========================================================================= + // ou_dec_complete 正确性测试 + // + // 流程: + // 1. 生成 BATCH 个明文,加密得到 FMLM 域密文 + // 2. 调用 ou_dec_complete 一次性还原明文 m(base-2^17) + // 3. D2H,写入 ou_dec_complete_test.txt,用 verify_ou_dec_complete.py 验证 + // ========================================================================= + { + const int BATCH = 200000; + + const size_t ct_bytes2 = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t m_bytes2 = (size_t)BATCH * OU_EXP_LIMBS * sizeof(uint64_t); + const size_t rp_bytes2 = (size_t)BATCH * OU_HR_EXP_LIMBS * sizeof(uint64_t); + const size_t tbl_bytes2 = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + const size_t texp_bytes2 = + (size_t)BATCH * OU_T_EXP_LIMBS * sizeof(uint64_t); + const size_t mout_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + + // ── t_64(ou_T base-2^17 → base-2^64)─────────────────────────────── + uint64_t t_64c[OU_T_EXP_LIMBS] = {0}; + for (int i = 0; i < ARR_LEN; i++) { + if (ou_T[i] == 0) continue; + int bit_offset = i * BASE_BITS; + int limb_idx = bit_offset / 64; + int bit_shift = bit_offset % 64; + if (limb_idx < OU_T_EXP_LIMBS) { + t_64c[limb_idx] |= (uint64_t)ou_T[i] << bit_shift; + if (bit_shift + BASE_BITS > 64 && limb_idx + 1 < OU_T_EXP_LIMBS) + t_64c[limb_idx + 1] |= (uint64_t)ou_T[i] >> (64 - bit_shift); + } + } + + // ── 生成明文 & 随机数 ──────────────────────────────────────────────── + srand(456u); + uint64_t *h_m5 = (uint64_t *)calloc(BATCH * OU_EXP_LIMBS, sizeof(uint64_t)); + uint64_t *h_rp5 = (uint64_t *)malloc(rp_bytes2); + for (int b = 0; b < BATCH; b++) { + uint64_t *m = h_m5 + (size_t)b * OU_EXP_LIMBS; + for (int j = 0; j < 21; j++) + m[j] = ((uint64_t)(rand() & 0x7FFF)) | + ((uint64_t)(rand() & 0x7FFF) << 15) | + ((uint64_t)(rand() & 0x7FFF) << 30) | + ((uint64_t)(rand() & 0x7FFF) << 45) | + ((uint64_t)(rand() & 0xF) << 60); + m[21] = (uint64_t)(rand() & 0x3FFFF); + } + ou_gen_r_prime(ou_p, BATCH, h_rp5); + + // ── 建表 & 加密 ────────────────────────────────────────────────────── + uint64_t *h_G_tbl5 = (uint64_t *)malloc(tbl_bytes2); + uint64_t *h_H_tbl5 = (uint64_t *)malloc(tbl_bytes2); + generate_G_table(Modn, ou_G, h_G_tbl5); + generate_H_table(Modn, ou_H, h_H_tbl5); + + uint64_t *d_G_tbl5 = nullptr, *d_H_tbl5 = nullptr; + uint64_t *d_m5 = nullptr, *d_rp5 = nullptr, *d_ct5 = nullptr; + CUDA_CHECK(cudaMalloc(&d_G_tbl5, tbl_bytes2)); + CUDA_CHECK(cudaMalloc(&d_H_tbl5, tbl_bytes2)); + CUDA_CHECK(cudaMalloc(&d_m5, m_bytes2)); + CUDA_CHECK(cudaMalloc(&d_rp5, rp_bytes2)); + CUDA_CHECK(cudaMalloc(&d_ct5, ct_bytes2)); + CUDA_CHECK( + cudaMemcpy(d_G_tbl5, h_G_tbl5, tbl_bytes2, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_H_tbl5, h_H_tbl5, tbl_bytes2, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_m5, h_m5, m_bytes2, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_rp5, h_rp5, rp_bytes2, cudaMemcpyHostToDevice)); + free(h_G_tbl5); + free(h_H_tbl5); + + ou_encrypt(d_G_tbl5, d_H_tbl5, d_m5, d_rp5, d_ct5, BATCH, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── H2D:密文 + 指数 t 计时 ────────────────────────────────────────── + uint64_t *h_ct5 = (uint64_t *)malloc(ct_bytes2); + uint64_t *h_t_exp5 = (uint64_t *)malloc(texp_bytes2); + CUDA_CHECK(cudaMemcpy(h_ct5, d_ct5, ct_bytes2, cudaMemcpyDeviceToHost)); + for (int b = 0; b < BATCH; b++) + for (int j = 0; j < OU_T_EXP_LIMBS; j++) + h_t_exp5[(size_t)b * OU_T_EXP_LIMBS + j] = t_64c[j]; + + uint64_t *d_ct5_in = nullptr, *d_t_exp5 = nullptr, *d_m_out5 = nullptr; + CUDA_CHECK(cudaMalloc(&d_ct5_in, ct_bytes2)); + CUDA_CHECK(cudaMalloc(&d_t_exp5, texp_bytes2)); + CUDA_CHECK(cudaMalloc(&d_m_out5, mout_bytes)); + + cudaEvent_t eh2d_s, eh2d_e; + CUDA_CHECK(cudaEventCreate(&eh2d_s)); + CUDA_CHECK(cudaEventCreate(&eh2d_e)); + CUDA_CHECK(cudaEventRecord(eh2d_s)); + CUDA_CHECK(cudaMemcpy(d_ct5_in, h_ct5, ct_bytes2, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_t_exp5, h_t_exp5, texp_bytes2, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(eh2d_e)); + CUDA_CHECK(cudaEventSynchronize(eh2d_e)); + float ms_h2d5 = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d5, eh2d_s, eh2d_e)); + CUDA_CHECK(cudaEventDestroy(eh2d_s)); + CUDA_CHECK(cudaEventDestroy(eh2d_e)); + + // ── 调用 ou_dec_complete ───────────────────────────────────────────── + float ms_comp = ou_dec_complete( + d_ct5_in, d_t_exp5, d_m_out5, BATCH, ou_p, ou_gp_inv, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, d_ctR, d_nctR); + + // ── D2H:明文结果计时 ───────────────────────────────────────────────── + uint64_t *h_m_out5 = (uint64_t *)malloc(mout_bytes); + cudaEvent_t ed2h_s, ed2h_e; + CUDA_CHECK(cudaEventCreate(&ed2h_s)); + CUDA_CHECK(cudaEventCreate(&ed2h_e)); + CUDA_CHECK(cudaEventRecord(ed2h_s)); + CUDA_CHECK( + cudaMemcpy(h_m_out5, d_m_out5, mout_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ed2h_e)); + CUDA_CHECK(cudaEventSynchronize(ed2h_e)); + float ms_d2h5 = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h5, ed2h_s, ed2h_e)); + CUDA_CHECK(cudaEventDestroy(ed2h_s)); + CUDA_CHECK(cudaEventDestroy(ed2h_e)); + + // ── 打印计时 ───────────────────────────────────────────────────────── + printf("\n========== ou_dec_complete test ==========\n"); + printf(" batch=%d ARR_LEN=%d\n", BATCH, ARR_LEN); + printf(" ------------------------------------------\n"); + printf(" [H2D ct+t] : %10.2f us (%8.2f us/op)\n", + ms_h2d5 * 1000.f, ms_h2d5 * 1000.f / BATCH); + printf(" [GPU dec_complete] : %10.2f us (%8.2f us/op)\n", + ms_comp * 1000.f, ms_comp * 1000.f / BATCH); + printf(" [D2H result] : %10.2f us (%8.2f us/op)\n", + ms_d2h5 * 1000.f, ms_d2h5 * 1000.f / BATCH); + printf(" ------------------------------------------\n"); + printf(" [Total] : %10.2f us (%8.2f us/op)\n", + (ms_h2d5 + ms_comp + ms_d2h5) * 1000.f, + (ms_h2d5 + ms_comp + ms_d2h5) * 1000.f / BATCH); + printf("==========================================\n\n"); + + // ── 打印前 3 条样本结果 ─────────────────────────────────────────────── + printf("[ou_dec_complete] 前 3 条 m_gpu(前 10 个 17-bit limb):\n"); + for (int b = 0; b < 3 && b < BATCH; b++) { + printf(" case %d:", b); + for (int j = 0; j < 10; j++) + printf(" %llu", (unsigned long long)h_m_out5[(size_t)b * ARR_LEN + j]); + printf(" ...\n"); + } + /* + // ── 写文件供 Python 验证 ───────────────────────────────────────────── + FILE *fp5 = fopen("ou_dec_complete_test.txt", "w"); + if (!fp5) { + fprintf(stderr, "[ERROR] 无法创建 ou_dec_complete_test.txt\n"); + } else { + fprintf(fp5, + "# ou_dec_complete test BATCH=%d ARR_LEN=%d BASE_BITS=%d\n", + BATCH, ARR_LEN, BASE_BITS); + fprintf(fp5, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp5, " %llu", (unsigned long long)Modn[j]); + fprintf(fp5, "\n"); + fprintf(fp5, "p:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp5, " %llu", (unsigned long long)ou_p[j]); + fprintf(fp5, "\n"); + fprintf(fp5, "T64:"); + for (int j = 0; j < OU_T_EXP_LIMBS; j++) + fprintf(fp5, " %llu", (unsigned long long)t_64c[j]); + fprintf(fp5, "\n"); + fprintf(fp5, "GP:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp5, " %llu", (unsigned long long)ou_gp_inv[j]); + fprintf(fp5, "\n"); + for (int b = 0; b < BATCH; b++) { + const uint64_t *m = h_m5 + (size_t)b * OU_EXP_LIMBS; + const uint64_t *ct = h_ct5 + (size_t)b * ARR_LEN; + const uint64_t *mdec = h_m_out5 + (size_t)b * ARR_LEN; + fprintf(fp5, "M:"); + for (int j = 0; j < OU_EXP_LIMBS; j++) + fprintf(fp5, " %llu", (unsigned long long)m[j]); + fprintf(fp5, "\n"); + fprintf(fp5, "CT:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp5, " %llu", (unsigned long long)ct[j]); + fprintf(fp5, "\n"); + fprintf(fp5, "MDEC:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp5, " %llu", (unsigned long long)mdec[j]); + fprintf(fp5, "\n"); + } + fclose(fp5); + printf("[ou_dec_complete test] 数据已写入 ou_dec_complete_test.txt\n"); + } + */ + // ── 释放 ───────────────────────────────────────────────────────────── + free(h_m5); + free(h_rp5); + free(h_ct5); + free(h_t_exp5); + free(h_m_out5); + cudaFree(d_m5); + cudaFree(d_rp5); + cudaFree(d_ct5); + cudaFree(d_G_tbl5); + cudaFree(d_H_tbl5); + cudaFree(d_ct5_in); + cudaFree(d_t_exp5); + cudaFree(d_m_out5); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_enc.cu b/heu/library/algorithms/ou_new/ou_enc.cu new file mode 100644 index 0000000..7e221e6 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_enc.cu @@ -0,0 +1,19678 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12420598007131, 404862574961, 1489476715458592911, + 16957267883830147813, 12043402377674848005, 16493683929612125450, + 6784216004999364667, 1572187082686061104, 12340520880335831418, + 9331149018221532047, 14282361491840612423, 8815005708360508971, + 3015263268809804842, 17865818108905781022, 9920159502644748834, + 16663441635849680526, 16754665591114057436, 18269507414440088358, + 10584069732924195264, 12369982964838682500, 17245146333661260496, + 14241460736918561983, 1495228790312205286, 650596912814140113, + 13561313259753582639, 12825917001300214751, 17648203335955277871, + 4937017841769177676, 8762411814459186696, 10916964928297372446, + 7934650143980081733, 16270299917683438672, 10599916200641161267, + 8153067147237791839, 3263540562937225549, 12427502011302527075, + 233468242765771385, 16905948202777045661, 7718378239207545448, + 12556027611679852532, 14355997036124831232, 10480393087867199263, + 8928529658741694940, 5329501958372121449, 14697840141414380795, + 4109255465868101899, 2189093870568573914, 2015298188623315682, + 15846288464350737625, 6154716281955634074, 14926805960051698119, + 6358221390244429391, 565479903706287841, 544001290941279643, + 8964697110262868308, 18143181760096837871, 4558614659708563008, + 13794757981490580748, 8464822194361423649, 3215767361281822704, + 3162014795756580172, 4046965963624353869, 13045127494050167872, + 2499196209129003472, 9910258488503570389, 17361620693833999278, + 9909125957613821241, 9160076386865281785, 1124681048989805109, + 13488748519780692092, 10062780796386632218, 1229663950817871070, + 9195972272372452954, 16317888560940382210, 7613588099474937282, + 4817436175622690809, 2044827062874516597, 10883598385799752225, + 15983746154563805830, 15305213912578498862, 15429027565390837948, + 15165749666894862326, 17094752132458668719, 17119153326423120334, + 16418978274070441262, 4020071962619695082, 1332089103620365082, + 2185942973674898504, 10623386207974578633, 17221670917682379713, + 1728235680214257322, 4351774997312820994, 13118214034020006923, + 17372706552005974065, 3635082452886395532, 9785476544995844000, + 5239966118419100751, 11446467251832755986, 8441993571518844572, + 17814051721047964692, 6635316643940291166, 591992533123407590, + 5392319256964754787, 6876233685503189066, 12560692609570806257, + 15922352178387947359, 10408865844002002467, 8334121862371121723, + 12325068481620777235, 16272399208182333772, 95913092914333903, + 9337511878088658138, 3620707693508746669, 2410461826666169902, + 15825609924068456517, 14631303029849624577, 12745725307310343054, + 11589389212418519430, 2402535340645687114, 10460197248544579201, + 2150167009874019400, 6579654444057751335, 10629733729325322988, + 15336435747150065712, 2033376745185295929, 5536738554015879544, + 6536459867767524988, 17547258140832520218, 16277632797518007476, + 17414093779241440205, 15172877332690969372, 11510709146567555164, + 4030987138388375538, 4545399202703904429, 3071928236698087396, + 6235823038589723214, 1385517152964008659, 5315971425862540039, + 737213825926788813, 12647942483396120777, 11729199013777796647, + 16701826799686034287, 10264009214043680734, 15712857482319384195, + 9683635541859048778, 12510245784932525987, 11453973523166319383, + 17641407202450751284, 7076031528159713278, 4418210339976245128, + 7042975543133840791, 2687978692681798677, 6620884024635649052, + 18356048427182277532, 12306609112475240772, 9532567709512160247, + 18330814574547253380, 1570356219924936203, 8589904627747755897, + 15801551039736301865, 14430949397822386754, 17457931947828226443, + 17066621315075642576, 5110941585241228080, 3113435720621532772, + 9109366487892874842, 6870609824263132877, 12909242032068451345, + 9978696359698762281, 2334245554206682016, 5474675901673095882, + 6791983755440567441, 7940389551176083495, 8857160764169799470, + 5079795210984134207, 13270967636848123637, 2817746455641616989, + 6823734589021892698, 6695712168139536841, 10011091548170035449, + 1880778261691426396, 4901098797777619034, 10275511334773938844, + 14807980597086261468, 2140936956496274302, 16031412187684387227, + 13312372442173260977, 3698443504872582944, 13391615949904799953, + 9911497060502159398, 6476394081735368123, 10884190273162182630, + 16812094898227487679, 17541511080948430118, 608781942072315334, + 11293244422430215336, 11785770688631267218, 1050397619038471432, + 8886815311739454498, 4768360092769081595, 6853419982646256117, + 4533330592021381799, 4065300292447750493, 5995762058657812513, + 5850499564985481696, 17494639705373259696, 5311542451880315192, + 10494352427384472421, 7053555067647955415, 13249062068286530876, + 4950095133067012336, 978416016041957203, 16132684490165566392, + 9461678038864020399, 13953140033946412303, 3128510702558997798, + 9989688919461798931, 7729804782429579018, 14657677073053593359, + 2033538150495053663, 18287519723785336581, 15945421906051418344, + 827199110667409676, 2705644704721997071, 7966095645914939971, + 16210256519322987223, 13500279664518420546, 6403250311096390047, + 15684661198519501341, 17612977561176850685, 1595673854820173645, + 11355330109781124768, 15152054633977714227, 7594126125750130131, + 1770764465613758168, 16103674937332423683, 12989244865424659639, + 985924609341537427, 17869081554936260320, 13756399484785420256, + 9417665840560873340, 6596619827292141536, 9282779747668579711, + 14087089070266996746, 16797541746528108966, 3649008588426316857, + 305957104569710731, 16307164143397236406, 8758406156783513520, + 8985939822111752351, 9823364479239409710, 14242788910385080985, + 16257804781363675711, 3460875886956782926, 5792587040463753957, + 8058808855494344097}; +const uint64_t con_modn_shoup[256] = { + 17258949916450502562, 8888112940397698219, 713580727021815326, + 7267352695482246207, 16882106641962718850, 4484957681626972705, + 11436540044429239973, 11422524016867871168, 14020805313791955351, + 3223598656225742213, 5973835375502394258, 18389866180389685123, + 16986950504460752843, 14151661301513655720, 17077263121249724353, + 9230648276690066133, 14180947083865490845, 9351212926919619729, + 14132401609129010813, 5094848903632022188, 9701476928191356694, + 2266122304659313991, 9642600553848677189, 17694947392610440444, + 4166068824698342286, 5461002667202655000, 12457224361386750838, + 17183384576469292838, 4577122507392282256, 212990517287292298, + 9593917155195774666, 16026775048324125596, 16942601617377608128, + 14517179166542822181, 5974534123485502056, 3098739027972382032, + 10256410151333955680, 17520563159797137486, 1437139722912674326, + 8992878724252091890, 12528748759234609070, 8869909966108752023, + 921545360850748754, 9072672431128383783, 5738738820990719155, + 14467747599114998535, 17288556119140904797, 10287232211619750677, + 14260418889536810543, 14236170919296794224, 3150044698371016975, + 3734472856964467879, 17367148054045259108, 8704684147659221983, + 2508292278448591736, 8252783960354095393, 5505177864158443900, + 6073691067055527962, 1797426869555915200, 10910566681039151494, + 7942925665011098265, 2975836443917988571, 12168486674970259490, + 16110198076303607574, 3176765343615554343, 16298611667673954183, + 17549972852929401068, 11301222901206006551, 16909965679872285631, + 4718899407226931290, 3249759906455460509, 12902044256019035476, + 2933343404927512513, 12635530155166610452, 1019795786854894305, + 5558319569785986415, 1736084326860448187, 1046201637077468719, + 8212969089240241408, 14337127378294037853, 386901202666220852, + 7007901116396992081, 12993354667310087947, 10362596985035361337, + 2632315220458003599, 15774429557071893885, 16402946062094934668, + 1455058537672538730, 10934206267514764178, 8653495955263034803, + 15626427345777561198, 9949509055664635150, 13551873570797002283, + 8464454269390886841, 17029555676738610993, 2576180826541960979, + 9277239235861014619, 5949389895883929521, 5961129371502370384, + 205015922583922578, 10998263302456627347, 630682958888475996, + 8272932531190805541, 11739326145926911850, 3711218700419121610, + 1995808726167056948, 9541493218673654357, 10947701895562230579, + 13307121742110758926, 5577478578055132785, 17291845771099354096, + 18386024476421024505, 12682871500298847381, 11130873186526484524, + 9511555001775281238, 12573445966195088218, 8735156148113491387, + 11986977738879871749, 3957140364191407682, 9447015894594054903, + 15075232292569908924, 3903378683888054983, 18108762427411217092, + 3848515571709448978, 9812945716174782266, 1439686790983251294, + 16833581092072042070, 16210195608156458739, 5586553670046771142, + 16285469091149528676, 9398773791615283024, 6409294867327805667, + 7451696515003176156, 9352800728836756410, 13006499844621233701, + 8943861815085160245, 9678575573847830426, 8255945162964001220, + 10730739327269562739, 5334149872806451563, 17598838495642584396, + 16369098332122517680, 17955871179034626630, 14641785369093985684, + 14862588269370659671, 18285743176298903762, 13037469746309732440, + 14680321993850845543, 4817891414280274891, 5149080315255804853, + 2128397328110922197, 10589038400063582404, 7361242319334319176, + 2633416174015188179, 8731722470221646364, 17695241939569347999, + 511281303145693578, 3915229765025655898, 12611201828639463913, + 12923773568272819998, 6599022632397876071, 13420679683499938789, + 2049055346114188916, 12969132418156712259, 4309319847942653387, + 18244990947493391329, 4321703802367122695, 6192941661180839271, + 3809220907288466934, 9686397595036380761, 5512567166054515334, + 7602079309929874744, 13912271487139773395, 14184657929002112052, + 14576671517535568717, 10645713420980679390, 6189820458973792135, + 11814662608890796689, 14130844520328466724, 5920079432947567926, + 11563681922574043697, 1737013972124835278, 9154718647461864771, + 4872088247015674989, 3298657241789078374, 4274525812376499266, + 2844739318930852390, 5423787114207982915, 711603148563500498, + 5237712164936082793, 16538115545248024540, 6726637234751909654, + 12914942507138411381, 13794083152765605836, 8092188286693750926, + 3585046709821797374, 9022641750310238763, 465420587122630117, + 13564884715383175347, 11129511873502417680, 5164151854448859979, + 15029903335193380621, 10868746496347154392, 4435651212450171864, + 13217684022763318278, 5933919845722166548, 3356965223531861545, + 5310913073599813054, 15054780849600338193, 6396900692144724485, + 7035542304377335696, 10193072637584812570, 13651162572970216008, + 3375517112149222934, 4439239668198625265, 4272637244444960409, + 12817026244019634782, 69354263255062928, 15302789662880427742, + 347862105615759955, 17769797205451930634, 9642432860176744960, + 12464642972155630307, 15664388251667466866, 11631517894848881873, + 17036907573636512589, 7095039048837509769, 17513417735648162476, + 17481830685724223605, 17280286146669482802, 2842388552256186635, + 3920667164872509047, 12274188700851890415, 10745285319594444214, + 1477783310694359069, 9823805546543362486, 13036068216114647306, + 2630587855422542122, 11133180172685530573, 5498501119950892568, + 12148543347986984583, 8699653759414714944, 12680167660387601338, + 7199244163763404941, 11286265673566327732, 16227882229914399460, + 13238315363757133297, 15029696251503760467, 3493558449506003864, + 1655355875307187670, 2509903151438279534, 16562697969364590150, + 16509371521602835055, 6068654425169380909, 12704430300130691212, + 1645898033077800833}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define OU_T_TAU 256 // 解密指数 t 的 bit 数(p-1 的大素因子,~256 bit) +#define OU_T_EXP_LIMBS (OU_T_TAU / 64) // = 4,每个 t 占 4 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 15 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +// 保留参数兼容调用处,128-bit 时不需要 p +static void ou_gen_r_prime(const uint64_t * /*ou_p_limbs17*/, int batch, + uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================================= +// ou_encrypt +// +// 功能:批量 OU 加密(正明文) +// c̃[i] = FMLM(G^{m_i}, H^{r'_i}) mod n (FMLM 域输出) +// +// 数学依据: +// OU 加密公式:c = G^m · H^r mod n +// FMLM 域语义:FMLM(Ã, B̃) = ÷B̃·R⁻¹ mod n +// 因为 Getgp/Getrn 输出 FMLM 域(乘了 R),所以: +// FMLM(G^m·R, H^r'·R) = G^m·H^r'·R mod n ✓ +// +// 三步流程: +// Step① Getgp :G 预计算表 × m_i → G^{m_i} (FMLM 域)→ d_ct_out +// Step② Getrn :H 预计算表 × r'_i → H^{r'_i}(FMLM 域)→ d_Hr(临时) +// Step③ XYfixWarpVector:d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) +// +// 调用前准备: +// generate_G_table(Modn, ou_G, h_G_table) 并上传 → d_G_table +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM +// 域,generate_G_table 输出) d_H_table [TABLE_SIZE × ARR_LEN] H +// 预计算表(FMLM 域,generate_H_table 输出) d_m_batch [batch × +// OU_EXP_LIMBS] 明文 m(base-2^64 小端序,OU_EXP_LIMBS=22) +// d_r_prime_batch [batch × OU_HR_EXP_LIMBS] 随机指数 r'(OU_HR_TAU=128 +// bit,OU_HR_EXP_LIMBS=2) d_ct_out [batch × ARR_LEN] 输出密文(FMLM +// 域) batch 明文数量 +// +// 返回:GPU 端 Step①②③ 总耗时(ms) +// ============================================================================= +float ou_encrypt(const uint64_t *d_G_table, const uint64_t *d_H_table, + const uint64_t *d_m_batch, // [batch × OU_EXP_LIMBS] + const uint64_t *d_r_prime_batch, // [batch × OU_HR_EXP_LIMBS] + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM 域) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getgp — G^{m_i} mod n(FMLM 域)→ d_ct_out + Getgp<<>>( + d_G_table, d_m_batch, OU_TAU, d_ct_out, batch, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step③:XYfixWarpVector — d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) + // 结果:c̃[i] = G^{m_i} · H^{r'_i} · R mod n(FMLM 域密文) + XYfixWarpVector<<>>( + d_ct_out, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================= +// ou_dec +// +// 功能:GPU 批量 OU 解密第一步——模幂 c^t mod n(普通域输出) +// +// OU 完整解密流程: +// Step①(本函数): c^t mod n ← FMLE_mod3_Kernel,NTT 参数为 n 模数 +// Step②(CPU 端): (result) mod p² ← 因 n=p²·q,c^t mod n 再 mod p² = c^t +// mod p² Step③(CPU 端): m = L(c') · gp_inv mod p,其中 L(x) = (x-1)/p +// +// 输入 d_c_tilde 为 FMLM 域密文 c̃ = c·R mod n; +// FMLE_mod3_Kernel 跳过 Step8(输入已是 FMLM 域), +// 保留 Step13(乘 R⁻¹,还原为普通域),输出 c^t mod n。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(FMLM 域,mod n) +// d_t_exp [batch × OU_T_EXP_LIMBS] 指数 t(压缩格式:OU_T_EXP_LIMBS=4 个 +// uint64_t, +// 小端序,位 i 在 limb[i/64] 的第 i%64 +// 位) +// d_output [batch × ARR_LEN] 输出:c^t mod n(普通域,base-2^17) +// batch 批大小 +// d_r0 [ARR_LEN] r₀ = (2^(ARR_LEN×BASE_BITS)-1) mod n +// (FMLM 单位元,即蒙哥马利域中的 1) +// 其余为 NTT 参数(n 模数,与 ou_encrypt 完全一致) +// +// 返回:GPU 端到端耗时(ms) +// ============================================================= +float ou_dec(const uint64_t *d_c_tilde, const uint64_t *d_t_exp, + uint64_t *d_output, int batch, const uint64_t *d_r0, + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * OU_T_EXP_LIMBS; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_t_exp + exp_off, OU_T_TAU, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +int main() { + uint64_t Modn[256] = { + 90459, 46625, 4442, 112182, 114968, 71534, 86507, 127438, 1685, + 52931, 104150, 120548, 110008, 7428, 47921, 17877, 51604, 95338, + 118828, 83898, 128532, 25064, 75321, 17612, 57426, 74847, 91485, + 17341, 32555, 39547, 13559, 24145, 41272, 116851, 19414, 130804, + 122030, 33140, 103122, 68029, 121047, 93139, 39899, 122526, 23084, + 120314, 105769, 34702, 77999, 9851, 111681, 49144, 6962, 87978, + 73280, 9230, 120490, 99393, 36052, 66751, 52136, 26210, 79015, + 53220, 88620, 106517, 78130, 39022, 74644, 83791, 78431, 76879, + 126508, 112379, 112791, 45016, 126862, 98555, 67445, 123550, 78839, + 94946, 54918, 116122, 88926, 128743, 56161, 106652, 89566, 6771, + 65824, 105202, 2315, 108169, 98258, 27489, 25367, 128358, 45812, + 97612, 65111, 13602, 17793, 42341, 58555, 122795, 20331, 7114, + 60301, 88265, 43599, 10109, 18770, 82809, 13834, 117670, 103794, + 49075, 68168, 87021, 105358, 120278, 55703, 7511, 51479, 98172, + 58294, 99843, 55022, 33556, 19683, 122106, 75656, 112874, 46791, + 112145, 106681, 122130, 87140, 32269, 9030, 59222, 342, 75944, + 94103, 130245, 53043, 71941, 64227, 216, 36445, 98379, 21169, + 62160, 91119, 86471, 3598, 23603, 98537, 23157, 124090, 112187, + 67586, 111947, 77293, 12685, 36844, 113775, 50048, 260, 118397, + 88108, 66275, 63944, 42517, 39246, 24220, 115910, 91147, 23953, + 45252, 7345, 75180, 16316, 123347, 38049, 74728, 78637, 21442, + 41317, 90948, 49684, 37151, 40998, 36225, 79231, 81693, 87993, + 55886, 2798, 82414, 110989, 110453, 57672, 14220, 82985, 1843, + 128765, 104271, 125420, 13546, 63807, 90504, 103167, 51426, 33297, + 102375, 75357, 118381, 38152, 11460, 114990, 16009, 63757, 78839, + 46610, 83493, 100684, 103041, 23340, 76918, 82210, 94466, 46728, + 76817, 44869, 84237, 127471, 98069, 60741, 890, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 89416, 26754, 112835, 104869, 63192, 109205, 3104, 29726, 113565, + 56327, 49094, 67751, 77782, 27488, 119631, 71590, 123612, 97756, + 84117, 11401, 55613, 66158, 70476, 101675, 91512, 52471, 103039, + 24137, 47894, 66258, 102515, 73290, 122151, 114889, 128986, 74711, + 74396, 24226, 116902, 16522, 70549, 97827, 20059, 59149, 52686, + 26411, 100945, 78310, 25925, 67014, 79121, 123619, 40826, 66607, + 30092, 89869, 105580, 20910, 93204, 44105, 73778, 21121, 113212, + 113354, 12462, 2188, 4706, 7519, 55812, 128963, 35815, 6339, + 118625, 95153, 80551, 71853, 110048, 116886, 28114, 36427, 99347, + 11391, 127843, 41917, 14312, 58727, 41417, 4720, 129585, 25994, + 106019, 13449, 40561, 107035, 2390, 33535, 7847, 22268, 12977, + 53724, 114048, 27127, 106539, 77244, 114474, 118543, 104524, 97508, + 116569, 73049, 95427, 14163, 96131, 121201, 90712, 20571, 129841, + 128480, 52654, 46075, 61521, 53523, 19041, 42853, 127248, 11120, + 123997, 130413, 56569, 98615, 56998, 99234, 71154, 41850, 65057, + 46995, 104286, 37295, 68580, 38759, 28487, 15348, 68337, 102177, + 78016, 31678, 52663, 1195, 70670, 125162, 16806, 115057, 38759, + 1825, 67883, 105361, 112649, 60917, 88939, 64087, 91926, 11035, + 53857, 46013, 61091, 124646, 128853, 106128, 66902, 108804, 8022, + 25699, 7098, 119054, 99103, 128189, 128398, 1319, 96584, 80228, + 81809, 54282, 6764, 98910, 13622, 81427, 9254, 125751, 67112, + 129121, 60904, 13975, 42521, 71104, 19357, 130027, 15821, 92416, + 81803, 63240, 43145, 90507, 34211, 103714, 62407, 43497, 11360, + 9542, 117006, 82808, 14980, 96632, 5649, 54610, 109476, 15174, + 113385, 48087, 34882, 126953, 67608, 7544, 126322, 75249, 4473, + 4549, 79315, 60258, 82932, 85897, 13892, 119911, 78557, 112723, + 15372, 51919, 45301, 73004, 122707, 121269, 108268, 19504, 49759, + 64648, 50866, 115280, 49268, 9806, 71075, 96540, 114884, 100204, + 79963, 109285, 92516, 38345}; + uint64_t R1[256] = { + 63840, 77643, 38589, 103335, 39211, 83771, 23554, 103602, 13011, + 33265, 99850, 107872, 65932, 9450, 98916, 44944, 70258, 117509, + 29436, 61780, 95269, 57764, 113636, 76719, 51738, 24730, 60149, + 46944, 29753, 55286, 121585, 3994, 70184, 56425, 18607, 115724, + 66839, 44507, 95968, 14278, 79412, 78754, 21044, 58743, 116286, + 99739, 26812, 85232, 130351, 42770, 22378, 65680, 11282, 130626, + 67800, 75341, 43729, 3643, 41124, 90487, 21576, 61026, 35178, + 84381, 65484, 94672, 34715, 116332, 18621, 103526, 117440, 57014, + 7488, 192, 2821, 11532, 82352, 92862, 129569, 103329, 115342, + 91522, 84549, 112318, 128186, 46600, 126196, 12657, 116, 12915, + 15556, 72466, 25692, 116854, 113547, 7800, 66899, 57704, 21718, + 49267, 129570, 121088, 11232, 26664, 76057, 98161, 88201, 110810, + 67711, 57350, 129067, 55009, 66719, 82831, 16394, 45675, 54590, + 128704, 53971, 119083, 90622, 38432, 89295, 6632, 106784, 29141, + 112848, 12094, 121287, 45606, 18268, 123439, 23933, 114383, 91280, + 32929, 75475, 33289, 81297, 101051, 117022, 93617, 78312, 23394, + 80435, 107542, 74032, 67728, 108966, 130175, 112943, 113217, 3110, + 3831, 11764, 90166, 19552, 130175, 21228, 105234, 12469, 127827, + 48250, 34863, 105366, 25850, 91309, 59492, 11089, 37416, 86082, + 43171, 97876, 17564, 124193, 22479, 73770, 129547, 35390, 7301, + 37318, 107441, 45540, 89189, 116809, 98763, 77827, 82906, 86323, + 5043, 40530, 35486, 24454, 64147, 22581, 127179, 28982, 25813, + 111759, 25875, 67144, 65925, 119945, 65987, 38832, 53596, 102445, + 109559, 69211, 82095, 14947, 117172, 24490, 3989, 530, 105261, + 20337, 23270, 104501, 40575, 111253, 100384, 110528, 20256, 52349, + 44932, 88144, 99082, 32981, 62102, 48256, 120608, 109895, 109835, + 5800, 110786, 116117, 15839, 97802, 39097, 461, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 7995, 85394, 119550, 88325, 48629, 56321, 47267, 70538, 124755, + 114206, 127949, 66234, 97879, 10116, 75346, 116315, 115411, 105906, + 104887, 101066, 124353, 105454, 127841, 24753, 117972, 83013, 23257, + 113424, 62775, 50508, 129544, 56423, 26208, 31226, 41958, 53567, + 44587, 99425, 44966, 34907, 58764, 54870, 85649, 49714, 43361, + 76738, 122439, 125767, 122331, 53573, 6701, 46414, 130886, 89782, + 95205, 67851, 93842, 96807, 57984, 70991, 49875, 67966, 100160, + 53735, 110509, 94681, 29919, 100667, 31561, 122128, 14547, 37662, + 81267, 127530, 117629, 62865, 45977, 80557, 102078, 7436, 34725, + 45975, 64707, 70771, 100145, 81892, 54893, 84306, 1885, 95921, + 84368, 11000, 119123, 35258, 123460, 71345, 38926, 122268, 62525, + 95283, 31822, 110123, 69419, 59439, 126346, 114333, 45092, 55125, + 41949, 99782, 5381, 92466, 79905, 99415, 99944, 3002, 115735, + 82477, 119373, 16392, 56646, 129157, 120981, 70091, 117510, 32163, + 77816, 39738, 79253, 31759, 34433, 38342, 64985, 89927, 47371, + 116134, 97731, 65556, 67396, 115184, 38875, 63741, 95486, 113931, + 113486, 14243, 9757, 94037, 117368, 14744, 97227, 47217, 89316, + 113910, 31862, 84729, 27111, 24920, 4721, 57424, 63294, 105420, + 70146, 45272, 113763, 128548, 130854, 38917, 94062, 130079, 79141, + 24590, 128415, 123042, 16028, 129234, 115403, 111824, 24030, 97272, + 98332, 72842, 105893, 113852, 20216, 73952, 51974, 34840, 63012, + 82468, 85351, 84814, 43057, 87894, 27393, 129833, 69731, 99955, + 34229, 102752, 81848, 58514, 8339, 56606, 33113, 89062, 78126, + 42786, 18705, 14075, 31097, 43319, 127767, 5814, 110148, 27445, + 4752, 28880, 96452, 37793, 26660, 4458, 96852, 84284, 28955, + 86794, 40007, 22025, 48903, 31630, 12571, 71804, 32197, 97091, + 16979, 65070, 112324, 13854, 102150, 46274, 56, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + // ========================================================================= + // ou_encrypt 正确性测试 + // + // ou_encrypt(m, r') = G^m · H^r' mod n(FMLM 域输出) + // + // 流程: + // 1. 建立 G/H 预计算表(CPU → GPU) + // 2. 随机生成 BATCH 个明文 m(base-2^64,OU_EXP_LIMBS=22 limb,正整数) + // 3. CPU 生成随机指数 r'(ou_gen_r_prime,128 bit) + // 4. H2D:m + r' + // 5. GPU:ou_encrypt → d_ct + // 6. D2H:密文 + // 7. 写入 ou_encrypt_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 1000; + + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t m_bytes = (size_t)BATCH * OU_EXP_LIMBS * sizeof(uint64_t); + const size_t rp_bytes = (size_t)BATCH * OU_HR_EXP_LIMBS * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + + srand((unsigned int)time(nullptr)); + + // ── Step 1:建立 G / H 预计算表(CPU),上传 GPU ────────────────────── + uint64_t *h_G_table = (uint64_t *)malloc(tbl_bytes); + uint64_t *h_H_table = (uint64_t *)malloc(tbl_bytes); + generate_G_table(Modn, ou_G, h_G_table); + generate_H_table(Modn, ou_H, h_H_table); + + uint64_t *d_G_table = nullptr, *d_H_table2 = nullptr; + CUDA_CHECK(cudaMalloc(&d_G_table, tbl_bytes)); + CUDA_CHECK(cudaMalloc(&d_H_table2, tbl_bytes)); + CUDA_CHECK( + cudaMemcpy(d_G_table, h_G_table, tbl_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_H_table2, h_H_table, tbl_bytes, cudaMemcpyHostToDevice)); + free(h_G_table); + free(h_H_table); + + // ── Step 2:随机生成 BATCH 个正明文 m(base-2^64,OU_EXP_LIMBS=22 limb)── + // 合法明文范围:|m| ≤ max_plaintext_ = 2^(BitCount(p/2)-1) + // p 最高位在 ou_p[80]=11(bit 1363),故 p/2 约 1363 bit, + // max_plaintext_ = 2^1362,取 m ∈ [0, 2^1362)。 + // + // 实现: + // limb 0~20(共 21 个):每个填满 64 bit → 21×64 = 1344 bit + // limb 21 :仅低 18 bit → 1344+18 = 1362 bit ≤ p/2 + // limb 22 及以上 :保持 0 + // + // 每个 64-bit limb 用 5 次 rand()(Windows 下 rand() 返回 15 bit)拼成: + // bits 0-14 : rand() + // bits 15-29 : rand() << 15 + // bits 30-44 : rand() << 30 + // bits 45-59 : rand() << 45 + // bits 60-63 : (rand() & 0xF) << 60 (取低 4 bit,补满 64 bit) + uint64_t *h_m = (uint64_t *)calloc(BATCH * OU_EXP_LIMBS, sizeof(uint64_t)); + for (int b = 0; b < BATCH; b++) { + uint64_t *m = h_m + (size_t)b * OU_EXP_LIMBS; + // limb 0~20:完整 64-bit 随机值(1344 bit) + for (int j = 0; j < 21; j++) + m[j] = ((uint64_t)(rand() & 0x7FFF)) | + ((uint64_t)(rand() & 0x7FFF) << 15) | + ((uint64_t)(rand() & 0x7FFF) << 30) | + ((uint64_t)(rand() & 0x7FFF) << 45) | + ((uint64_t)(rand() & 0xF) << 60); + // limb 21:低 18 bit(确保 m < 2^1362 < p/2) + m[21] = (uint64_t)(rand() & 0x3FFFF); // [0, 2^18) + } + + // ── Step 3:CPU 生成随机指数 r'(128 bit)────────────────────────────── + uint64_t *h_rp = (uint64_t *)malloc(rp_bytes); + ou_gen_r_prime(ou_p, BATCH, h_rp); + + // ── Step 4:H2D,计时 ───────────────────────────────────────────────── + uint64_t *d_m = nullptr, *d_rp = nullptr, *d_ct = nullptr; + CUDA_CHECK(cudaMalloc(&d_m, m_bytes)); + CUDA_CHECK(cudaMalloc(&d_rp, rp_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(d_m, h_m, m_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_rp, h_rp, rp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 5:GPU ou_encrypt,计时 ───────────────────────────────────── + float ms_gpu = ou_encrypt( + d_G_table, d_H_table2, d_m, d_rp, d_ct, BATCH, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 6:D2H,计时 ───────────────────────────────────────────────── + uint64_t *h_ct = (uint64_t *)malloc(ct_bytes); + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(h_ct, d_ct, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 7:打印计时 ────────────────────────────────────────────────── + printf("\n========== ou_encrypt test ==========\n"); + printf(" batch=%d ARR_LEN=%d OU_TAU=%d OU_HR_TAU=%d\n", BATCH, ARR_LEN, + OU_TAU, OU_HR_TAU); + printf(" [H2D m+r'] : %8.4f ms (%6.4f ms/op)\n", ms_h2d, + ms_h2d / BATCH); + printf(" [GPU ou_encrypt] : %8.4f ms (%6.4f ms/op)\n", ms_gpu, + ms_gpu / BATCH); + printf(" [D2H ct] : %8.4f ms (%6.4f ms/op)\n", ms_d2h, + ms_d2h / BATCH); + printf("=====================================\n\n"); + /* + // ── Step 8:写测试数据文件 ──────────────────────────────────────────── + // 格式 ou_encrypt_test.txt: + // 第 1 行(注释):元信息 + // 第 2 行 N: OU 模数 n(256 limb,base-2^17) + // 第 3 行 G: 公钥 G (256 limb,base-2^17) + // 第 4 行 H: 公钥 H (256 limb,base-2^17) + // 每用例 3 行: + // M: 明文 m (OU_EXP_LIMBS=22 limb,base-2^64) + // RP: 随机数 r'(OU_HR_EXP_LIMBS=2 limb,base-2^64,OU_HR_TAU=128 + bit) + // CT: 密文 c̃ (256 limb,base-2^17,FMLM 域) + FILE *fp = fopen("ou_encrypt_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_encrypt_test.txt\n"); + } else { + fprintf(fp, + "# ou_encrypt test BATCH=%d ARR_LEN=%d BASE_BITS=%d" + " OU_TAU=%d OU_EXP_LIMBS=%d OU_HR_TAU=%d OU_HR_EXP_LIMBS=%d\n", + BATCH, ARR_LEN, BASE_BITS, OU_TAU, OU_EXP_LIMBS, OU_HR_TAU, + OU_HR_EXP_LIMBS); + // N: modulus n + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + // G: public key G + fprintf(fp, "G:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_G[j]); + fprintf(fp, "\n"); + // H: public key H + fprintf(fp, "H:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_H[j]); + fprintf(fp, "\n"); + // 逐用例写 M / RP / CT + for (int b = 0; b < BATCH; b++) { + const uint64_t *m = h_m + (size_t)b * OU_EXP_LIMBS; + const uint64_t *rp = h_rp + (size_t)b * OU_HR_EXP_LIMBS; + const uint64_t *ct = h_ct + (size_t)b * ARR_LEN; + // M: 明文(base-2^64,OU_EXP_LIMBS=22 limb) + fprintf(fp, "M:"); + for (int j = 0; j < OU_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)m[j]); + fprintf(fp, "\n"); + // RP: 随机指数 r'(base-2^64,OU_HR_EXP_LIMBS=2 limb) + fprintf(fp, "RP:"); + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)rp[j]); + fprintf(fp, "\n"); + // CT: 密文(base-2^17,ARR_LEN=256 limb,FMLM 域) + fprintf(fp, "CT:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ct[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_encrypt test] 测试数据已写入 ou_encrypt_test.txt\n"); + } + */ + // ── 释放 ────────────────────────────────────────────────────────────── + free(h_m); + free(h_rp); + free(h_ct); + cudaFree(d_m); + cudaFree(d_rp); + cudaFree(d_ct); + cudaFree(d_G_table); + cudaFree(d_H_table2); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_mulplain.cu b/heu/library/algorithms/ou_new/ou_mulplain.cu new file mode 100644 index 0000000..86f2755 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_mulplain.cu @@ -0,0 +1,19145 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 11687101653036, 18446743758138104250, 14211471509222940197, + 4235272491137479070, 18240298465959804861, 3612012136679214338, + 17250142769850486366, 16237778209286303069, 4059562170797630074, + 9869391574692501418, 11541833079389403685, 6455235399891900791, + 12385636973677184651, 16642176613718519917, 18371701685023660666, + 12908184103137142305, 432025558415355680, 13421505415041197067, + 3981525027614746084, 11947396656646251748, 17803480186391939853, + 8964175857302805981, 10755671907831908778, 1904820253353466484, + 10527642910445979147, 13113465356235953778, 7234975774234226545, + 9757071334591424272, 7316737729072523157, 6883729646458586186, + 8932421558757596044, 14597307983516136610, 14233691893978885667, + 11827243429966767869, 4500161214914298038, 10200789522299270510, + 13311736198424319672, 15003914107128413955, 15970198162388647651, + 5178000660069144055, 12257227048244175003, 199311466455739912, + 1199638074941611369, 7733792994982443122, 5589666049742506788, + 13804186403958915190, 17844357954872068228, 1608648031291287388, + 14833453796159433687, 17508515457688533059, 5757945642895465137, + 13081582882382389324, 5394527028006066918, 276092650195297071, + 17347510335080686092, 1701269284563161833, 13303804442297711418, + 8253121806998843455, 1803714610749342533, 15051344875346831329, + 17235261944528818002, 8632334691816709691, 7259437303239191782, + 15692170915486480673, 16909097158836466193, 3813100579756643340, + 8120672335331207645, 17658082942068012540, 14712625527555672008, + 3013490167685507391, 17053781224112072993, 7951833678156564288, + 16882129623470333590, 16598623833219974769, 8844055318436626140, + 7109452390029814739, 17973202994978822509, 17988667507327805934, + 4123463954687411915, 4940157814071223716, 10034667236662981698, + 2823422309362454166, 8801094461519448259, 16352417325266577247, + 4120137837347666316, 13914619055383639681, 13924524887541743843, + 16792129803426718801, 17665685411675475587, 4424050859418317508, + 7334727500762095896, 17503735803347907081, 17481354675564559188, + 3960789173235012097, 2382836270529746228, 4344431928835757414, + 10028103253977297380, 10742172769839739841, 4046021704283408275, + 15048763124930553406, 12884462872804057858, 6406919215710243952, + 2305936473049040111, 13121402678735537817, 4908133312471037493, + 11377681924055600462, 12660562847764485932, 4317962794267615766, + 1681209879049333151, 17555975506207345043, 5125753427319466102, + 3335447880892075219, 8356915887911857374, 8584812879987417990, + 17895049134452250854, 9085212301170466888, 7267866673654176139, + 11010976396069557933, 3608178248288276404, 676841772753514930, + 14867830803115014463, 1834280874555657868, 16310358636794094835, + 7330665673582977596, 15741791681143414831, 210676699798228406, + 17198551982049102727, 6109711417879930926, 10546103406410231004, + 15078747006884146927, 15249364729241398859, 7659463200845052268, + 6442795927427660859, 2250605931405808055, 8092318475578226474, + 18259756362431830946, 17518863421902388405, 10430337473484617894, + 11467857499548239759, 14850024957617392139, 4520997378651548347, + 1002071158803202119, 13705130616563662128, 16248739479905188442, + 3181222670542328118, 10465632553806001024, 13994891411389346011, + 10984874395183979750, 9226503659276477277, 16804196055515371214, + 10159135197864231902, 2843438330111370559, 10398977526462394360, + 13629959477406403811, 17564676539751491213, 9169917922917002660, + 9102739415085474595, 2571556159645270064, 15524480249380336561, + 16515752187799014418, 15631770625703612855, 15987278742587054296, + 6287084574542908707, 5110496278831165043, 3153708236541173992, + 12927769102029794814, 15247894309294568262, 9307521059673973398, + 2770646931232435574, 16454025378125412919, 10633343977751814089, + 1784373777963240394, 6934373261888962825, 17349218505508302052, + 10968015619489026254, 1752595997967202555, 12837659792768865327, + 14742784040235182007, 10962127976630829136, 5900158086628228693, + 17316277940499271723, 348494413055888541, 488358709608098659, + 10382164829144562111, 10796492275262689178, 10517686393808888673, + 779545184377220173, 3381663347793212032, 15236282164919686367, + 6334307549334192514, 3063003522052060686, 12114810039050492787, + 8870556826759400006, 4038453701720268007, 14314379071943608840, + 4339980657355467038, 9171890896160995321, 7917821014549284449, + 13180571956635383946, 18248798750186441796, 14577763713235528379, + 1809799949882400184, 6318379507872769143, 1709904138639204836, + 4433595655137503557, 14198791200961044293, 4650959752702308354, + 3318171583106947229, 209694021303727738, 1076839995989163599, + 16905707527716642672, 3695319746732430913, 10252674333030518371, + 16420700828256874086, 2433845634309674122, 16595131843996099370, + 16829576163922136336, 4841410023332049560, 12592434475652961886, + 17572224200405592702, 2431385938003789647, 10061028979483934059, + 2925075822122586101, 8606434114160323337, 5607119374066730577, + 11884780541782053893, 3126131661631420656, 10027052590555524485, + 16048091853305621015, 15852396680435263215, 17385109108245871297, + 12899005442699559936, 7848549015331456901, 9096729734807002481, + 4996467004929486051, 12243245730936161727, 9057745574396783344, + 7771655314147204603, 17881871823990369590, 15212325419875966733, + 11042754214829301385, 3281380445998958437, 17088850123667971831, + 14495632125498788619, 17994272834268936450, 16829150404837372, + 6402480299793703941, 10533393325012975763, 2416806924432625423, + 2875142845742952402, 12269921477737603466, 712826757316093295, + 4415075740707273176, 7975119161839365838, 11673666813117204770, + 9840315584201792241}; + +const uint64_t con_modn_shoup[256] = { + 14252880640204352951, 18322338132012530560, 6849211572557864124, + 15932334426407894751, 14177538223384229049, 9106170071204927292, + 16758578487033067404, 5864117390783715525, 16494545899140440021, + 5897258617717822505, 1933174352700629569, 12810083258791009448, + 12690514865985841899, 3970720354745169798, 4239814533413767079, + 5609102486863112468, 11230284723594595426, 17034417004294615591, + 5986132948557996967, 1868566874544028188, 2158239585541928173, + 17841097290863850509, 9305374060203424222, 10083694270531949160, + 2654649954734757684, 17721823101598855865, 599980504891318557, + 17732121566266018424, 1832524248816892725, 5295674783104160211, + 13283213776815003617, 10900691717351424196, 6021057974446928650, + 12624795618036119368, 2798162278377969124, 7399862538297940480, + 576839721297220460, 16060704215571397670, 16380205270440154947, + 7499979448419237887, 13254841413012858481, 10664669973443882596, + 1765312882701300737, 9426266066319221551, 12823009007753704482, + 10630699336822577567, 16298120910453621338, 13674950148586695572, + 17678273253225120972, 6798806775207395115, 13410427498759750653, + 2614783784077562964, 9342414901595102647, 14373786281595714575, + 6330183866169305354, 998675938268783033, 10221732541776071598, + 8979762078911881033, 9878667596621300249, 4856285279936479033, + 14833025980849776526, 5604878655902262946, 13088552648421720803, + 1801154013367199521, 4158119999782205988, 7904891504652660262, + 9042945763108429841, 4642264688478488779, 16204979912313018920, + 3580705517878336362, 9712754433060271621, 8675179099278674516, + 8186655178929728093, 1884659203003867161, 17775229374523385263, + 6390348527000753038, 6439058351892770174, 2339745637453507323, + 16274314407512660647, 2247518004490005028, 18003796786185432156, + 4540940947376355923, 11987538437574474975, 16166798012420960901, + 16121611900792272328, 16082928115738878740, 121528093926685229, + 11609994994605905995, 2593441955413327993, 16920803883743198476, + 11945409668615507125, 15459882499135165139, 4709903422099897132, + 1915945056478813527, 17487099108173624447, 4121351438621439846, + 11648490996845515622, 16906896413860707859, 440932069689474224, + 5596373384320545758, 6286719224840257488, 15070666469307485122, + 1718056780076659255, 16292491121877970301, 16399246121914763003, + 868264559834958645, 5880650461523548368, 13037697811177873232, + 12598280349103069353, 8787439026840426841, 75102682845531848, + 10793543124523682506, 4058772666671965704, 4575391113880810276, + 7977675084418792789, 3637051392050280908, 16362683407568863478, + 18347383388798481077, 9115743514592553391, 9421569851894468249, + 15594101773322529942, 11807267208082355523, 2211845703086696074, + 17348335706114235958, 11847926503719721254, 17547278040398801999, + 5869056178242350580, 13320003773654588467, 12699478824066277151, + 5239882100070320334, 7261256595809982529, 18328110662323898796, + 14262563528151153433, 10570694000294503318, 12813828885908106934, + 10929763919809758798, 2938308820079533814, 12010181661893483546, + 5724066348617601097, 14693589767406584560, 7346156359909105313, + 12463844683080763521, 6157213913132141689, 10056871538135474507, + 10533920527198962255, 8235152268042627670, 12319030087627737531, + 8540756947882872704, 7431325835456426550, 1301653294393655025, + 5378735902526386334, 14912383613060771114, 639721130790109100, + 16570337161183830109, 3674985562098081097, 25882515425358888, + 5781372063417524209, 16445334884700810773, 3544553957273777187, + 7670642182104980993, 9626872654279745485, 16105190074590295359, + 7770490841006776992, 15228876210409060859, 9849662374993193055, + 8654391918266496929, 6489400423788825431, 3633531112627186925, + 14858949636521671041, 1105232854426717343, 16217154593325743207, + 6106199821319202195, 15821125396653981131, 397434568115144159, + 5468761408652955936, 1296217405573136392, 8004677824854586354, + 2227606275875858042, 810603102045699119, 10604814613007946604, + 5290458938805336352, 17851909192937068431, 13268718299834334195, + 10806219279687765004, 7326952643401977865, 5984244256621982617, + 1659224770285885321, 10142490388121661620, 543184966312697096, + 6161334213393132777, 8606758596526178508, 16789552215120228336, + 8023727110355433065, 4647377966945102884, 6753109868947696463, + 3601294586865406454, 17607120286020078979, 2828322973879047780, + 2912791741152772408, 12563760677479960004, 14534280132691080513, + 1458078075305867271, 965960906590360464, 7718560025107401880, + 5982496231126478023, 9871084626240060187, 13176440103626612310, + 12254705932735020033, 6091017959275360856, 8575903195682037371, + 12248661790925770989, 15874428453561837902, 11580211822667122432, + 1675684581791044528, 953808119473653258, 5212502010992923248, + 11653707338854681558, 2814312282799796103, 10741977896772006037, + 792456711582471131, 5829394638712031862, 14582582386261703791, + 15116195068952376181, 16192690152961596921, 14996186982344279167, + 12452747948715197047, 9822110408961686525, 13084672463213903010, + 15511412028972873361, 4378034570765898628, 7434337709561426193, + 10855303285220731112, 7759227917166922418, 4976939851858003292, + 5204453497818198107, 16768838491807377833, 9561958848337334713, + 703798444570847189, 14816217796224051499, 2968875028840941609, + 6519175664834707328, 17450997194073458375, 12811758118865484166, + 10759621678827047865, 11859099579116844022, 14425180740705110568, + 7511257586748720299, 3539736587530686592, 4447312216206097890, + 11184913710261542236, 8771977661991653642, 12354338902272316516, + 15541265547581959679, 14587017360825887711, 15248329331280087903, + 10992628732632760992}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +int main() { + uint64_t Modn[256] = { + 43351, 84159, 17007, 126963, 115814, 64975, 14865, 122878, 58093, + 76773, 16638, 87086, 110462, 105466, 35053, 36095, 8051, 116177, + 119699, 118157, 12357, 71314, 68424, 35266, 58013, 63468, 22117, + 10903, 124058, 90359, 68490, 117774, 56449, 45990, 26837, 86153, + 120741, 31603, 78596, 24019, 45134, 33649, 61458, 59406, 88868, + 60745, 113313, 123484, 30017, 98185, 93108, 73040, 39521, 18181, + 2647, 51647, 10194, 73702, 22934, 64, 29664, 94536, 9414, + 63827, 6028, 107137, 71399, 49216, 8196, 46100, 117329, 67195, + 25041, 122567, 110161, 82524, 85064, 85420, 38367, 90728, 6216, + 87366, 124652, 29067, 100922, 38894, 64688, 22860, 83774, 130371, + 39036, 94816, 45277, 76221, 67984, 78245, 70889, 64430, 52640, + 50933, 54580, 32496, 95587, 110988, 102834, 68631, 42744, 111149, + 127114, 116295, 108662, 4710, 31837, 15424, 50234, 99229, 61393, + 81585, 33195, 14128, 9168, 55047, 119038, 97329, 43164, 111637, + 39396, 13009, 90209, 92184, 81272, 101938, 57149, 82121, 100630, + 37780, 7881, 13181, 8505, 125111, 43862, 119168, 19431, 80034, + 114187, 71294, 52911, 81495, 14533, 87246, 126978, 30310, 9978, + 44551, 60081, 126942, 75376, 77030, 36034, 104993, 58885, 90371, + 111023, 45378, 97203, 126393, 72942, 8192, 124336, 37338, 116797, + 66693, 60337, 12040, 90738, 108119, 66171, 78981, 79494, 91989, + 89494, 118041, 29798, 30883, 110522, 122729, 7823, 62523, 20666, + 52089, 43045, 51146, 24317, 38753, 122735, 100047, 56716, 117911, + 60032, 27220, 44093, 56, 113553, 49629, 84418, 64845, 67097, + 10050, 8296, 90055, 75973, 63190, 82919, 56713, 30800, 90227, + 63208, 39501, 61899, 129744, 78395, 58460, 121961, 72489, 20054, + 64673, 102069, 663, 109348, 36701, 18676, 98666, 77634, 108307, + 103092, 49932, 49372, 39729, 72337, 90146, 247, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49877, 103520, 14633, 105196, 88121, 75472, 97344, 77893, 104350, + 114542, 91892, 50952, 64070, 21857, 78639, 2966, 6840, 35973, + 24928, 25085, 49586, 41185, 13904, 75536, 2710, 53110, 104559, + 109441, 40809, 28751, 119672, 47293, 31768, 78934, 55525, 52600, + 69445, 27490, 96611, 84604, 11133, 106316, 35838, 34416, 127179, + 34651, 118161, 16314, 45331, 107582, 24235, 12592, 81941, 19278, + 85161, 104589, 13239, 111883, 81763, 64318, 45287, 80867, 37986, + 126633, 118333, 32371, 49588, 111307, 6225, 98536, 50417, 128492, + 7707, 81120, 51703, 13989, 14502, 86184, 74055, 95503, 80070, + 61934, 73173, 322, 128915, 100622, 92460, 3170, 90102, 44305, + 79192, 110628, 84896, 92155, 13698, 129632, 82311, 123141, 99527, + 28216, 52443, 78019, 37062, 11803, 15622, 2177, 42554, 81945, + 37634, 97471, 11261, 30170, 62893, 129991, 77778, 123677, 75667, + 22518, 67693, 79122, 34634, 117067, 33393, 41664, 46880, 63508, + 1202, 120484, 35970, 14884, 64323, 87199, 61041, 17700, 6496, + 128561, 104072, 129267, 31935, 119451, 54205, 55539, 59285, 49953, + 34036, 45769, 28494, 30664, 79616, 47756, 57833, 99317, 122074, + 130764, 65477, 42562, 43963, 97923, 114657, 34735, 19193, 84201, + 24821, 98149, 101368, 100860, 110862, 96240, 101375, 55675, 99994, + 12323, 56026, 120364, 84030, 10327, 108568, 4795, 122128, 3775, + 1740, 113334, 5740, 61052, 15255, 44939, 84950, 4631, 87197, + 63464, 45041, 49844, 102052, 41710, 76594, 96748, 2213, 15033, + 56862, 42121, 22702, 29104, 55955, 32193, 62378, 61812, 37549, + 27929, 118796, 116386, 35884, 83278, 116744, 103768, 106752, 29801, + 6976, 81713, 55669, 12038, 51733, 6915, 128541, 82038, 102167, + 64630, 125581, 69829, 79662, 80895, 89416, 41571, 113918, 73736, + 22655, 72892, 97009, 75512, 83469, 50798, 35893, 72631, 27752, + 114176, 116066, 35199, 11556, 117400, 53979, 71662, 76589, 35790, + 51797, 38295, 48839, 44050}; + uint64_t R1[256] = { + 95787, 87855, 74590, 120089, 96462, 125333, 89873, 62820, 56744, + 93675, 114260, 58407, 55044, 4742, 20922, 129032, 18634, 103071, + 2852, 114517, 116272, 79216, 95365, 35177, 128432, 80425, 18923, + 592, 13588, 42144, 48019, 39668, 66805, 42663, 33194, 65911, + 93428, 9610, 76041, 48300, 121686, 67062, 30099, 86626, 99273, + 90908, 66468, 28073, 71719, 43868, 40340, 88274, 109318, 21824, + 16472, 116161, 1346, 106033, 20342, 35258, 20632, 105594, 118266, + 97653, 97643, 7306, 67863, 79950, 112151, 117205, 39906, 100559, + 19328, 14826, 43881, 64539, 123341, 37113, 15909, 43631, 27755, + 12868, 87791, 110907, 2763, 41576, 76238, 21079, 28709, 67173, + 22692, 45867, 80137, 111080, 90017, 19215, 15056, 7843, 34411, + 10397, 47157, 113197, 77959, 43337, 123310, 26898, 103324, 95568, + 37773, 53742, 58532, 64900, 12429, 109482, 75505, 70429, 89935, + 67404, 103144, 45403, 28839, 100826, 80183, 60279, 60825, 67114, + 15456, 95163, 5820, 106812, 38605, 127798, 43023, 23037, 109334, + 82354, 36764, 29882, 1460, 109709, 70002, 61938, 129339, 12574, + 23578, 116415, 67219, 51854, 56951, 3482, 98043, 101818, 116934, + 91679, 101444, 12152, 24202, 116763, 23200, 65030, 72892, 11401, + 89306, 104758, 94473, 69024, 2331, 120575, 38998, 71116, 123520, + 12246, 31511, 15417, 54824, 35449, 51579, 129451, 119392, 117118, + 124644, 45532, 83498, 101978, 4264, 7165, 33903, 43033, 21694, + 89206, 4834, 104846, 13789, 111463, 68368, 35303, 42736, 93343, + 26110, 67668, 100827, 18175, 80987, 95026, 43049, 29875, 73340, + 103134, 127164, 12263, 68500, 44986, 40419, 118430, 83125, 27361, + 83899, 599, 102652, 96223, 98631, 95859, 60979, 40494, 118153, + 71598, 113945, 125768, 2378, 115813, 13443, 7038, 3598, 53377, + 26625, 110573, 3294, 119847, 72218, 85832, 125, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 2023, 29313, 55760, 30415, 117209, 68859, 126088, 105595, 44062, + 129614, 108803, 62812, 119884, 130526, 4554, 108683, 53422, 114202, + 34719, 66446, 65067, 23670, 9591, 82680, 20040, 77609, 34903, + 85401, 37760, 43899, 14395, 67595, 12710, 101806, 77200, 103316, + 65099, 125953, 105217, 67808, 5979, 63677, 12355, 111674, 70131, + 45594, 117232, 10425, 81356, 112691, 14758, 10490, 24556, 27922, + 32980, 20928, 118420, 17204, 4244, 126937, 116849, 106497, 51321, + 114935, 45503, 15461, 59271, 111583, 30113, 103352, 10622, 32510, + 41116, 86928, 10137, 101567, 30707, 124356, 108755, 8203, 11158, + 9603, 114740, 5093, 13054, 61800, 75687, 38080, 11550, 87153, + 33247, 18929, 66437, 13511, 39575, 18765, 61155, 77315, 112366, + 76906, 125693, 40793, 40582, 43161, 81338, 111531, 84813, 49322, + 71309, 83250, 14948, 44745, 13967, 98243, 116072, 5842, 82567, + 77993, 80649, 107659, 66320, 122438, 54394, 54983, 79006, 105681, + 94840, 79085, 41950, 106863, 130420, 89427, 83726, 86511, 44750, + 12837, 47751, 33678, 115313, 66053, 43941, 100068, 107956, 169, + 62673, 70167, 106884, 28460, 37125, 129538, 93387, 56010, 35558, + 40621, 77463, 39765, 16013, 85203, 39465, 19946, 122865, 58068, + 60861, 54036, 48234, 130529, 59321, 83170, 116672, 11733, 94357, + 35207, 46800, 22992, 47306, 80973, 28136, 59828, 4338, 109019, + 16604, 58136, 123247, 75151, 11948, 105570, 86992, 72951, 15873, + 12088, 27387, 60538, 80333, 72819, 97245, 41421, 126078, 82127, + 78747, 70199, 99633, 38926, 56088, 20821, 130509, 125699, 127559, + 47166, 102354, 71946, 79468, 81253, 11120, 37016, 20384, 95174, + 126298, 33756, 60551, 106592, 43758, 59029, 22997, 63753, 65367, + 47855, 5100, 118402, 102326, 108278, 49481, 99388, 114710, 101475, + 124482, 27406, 78309, 52397, 69291, 787, 10, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ========================================================================= + // ou_mulplain 正确性测试 + // + // 公式:c̃^k → c^k·R mod n (蒙哥马利域全程,OU 密文乘明文) + // + // 流程: + // 1. 随机生成 BATCH 个密文(base-2^17,< n) + // 和明文(base-2^64,< p 且接近 1363-bit) + // 2. H2D 传输密文 + 明文,用 cudaEvent 计时 + // 3. 调用 ou_mulplain,内核计时由返回值给出 + // 4. D2H 传输结果,计时 + // 5. 写入 ou_mulplain_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 200000; + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)BATCH * OU_EXP_LIMBS * sizeof(uint64_t); + + // 私钥 p 的 base-2^64 表示(22 limb,LSB-first) + static const uint64_t p_limbs[OU_EXP_LIMBS] = { + 97859824528822159ULL, 1819502311735909913ULL, + 7447557854589657083ULL, 785921163302348405ULL, + 1559236396946863557ULL, 17984350884635806642ULL, + 12716701552721899659ULL, 12882653723943309200ULL, + 7080848173246890028ULL, 12362263030800919216ULL, + 5829044229350676221ULL, 10697051272769895195ULL, + 12041503530722981302ULL, 17803059532771415190ULL, + 9187126869447618656ULL, 15360160842400019836ULL, + 16286525798957055928ULL, 2276504101893439303ULL, + 9897714879286757503ULL, 15224956524078906016ULL, + 17667842374690572841ULL, 485219ULL}; + // p 最高 limb = 485219;k[21] ∈ [2^18, 485218] 保证 k 为 1363-bit 且 < p + const uint64_t P_TOP = 485219ULL; + const uint64_t K_TOP_MIN = 262144ULL; // 2^18 + + // RAND64:5 次 15-bit rand() 拼出 64-bit(Windows RAND_MAX=32767 安全) +#define RAND64() \ + (((uint64_t)(rand() & 0x7FFF) << 49) | ((uint64_t)(rand() & 0x7FFF) << 34) | \ + ((uint64_t)(rand() & 0x7FFF) << 19) | ((uint64_t)(rand() & 0x7FFF) << 4) | \ + ((uint64_t)(rand() & 0x000F))) + + // ── Step 1:随机生成密文和明文 ─────────────────────────────────────────── + const uint64_t MASK17_T = (1ULL << BASE_BITS) - 1ULL; + const uint64_t CT_TOP = Modn[240]; // = 247,密文最高 limb 严格上界 + + uint64_t *h_ciphers = (uint64_t *)malloc(ct_bytes); + uint64_t *h_plains = (uint64_t *)malloc(exp_bytes); + uint64_t *h_results = (uint64_t *)malloc(ct_bytes); + + srand((unsigned int)time(nullptr)); + + // 密文:256 limb base-2^17,limb[240] < CT_TOP,其余高位 = 0,确保值 < n + for (int b = 0; b < BATCH; b++) { + uint64_t *c = h_ciphers + (size_t)b * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < 240) + c[j] = (((uint64_t)(uint32_t)rand()) ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + else if (j == 240) + c[j] = (uint64_t)rand() % CT_TOP; + else + c[j] = 0ULL; + } + } + + // 明文:22 limb base-2^64 + // k[0..20] 随机 64-bit;k[21] ∈ [K_TOP_MIN, P_TOP-1] + // → k ≥ 2^1362(1363-bit)且 k[21] < P_TOP → k < p + for (int b = 0; b < BATCH; b++) { + uint64_t *k = h_plains + (size_t)b * OU_EXP_LIMBS; + for (int j = 0; j < OU_EXP_LIMBS - 1; j++) k[j] = RAND64(); + k[OU_EXP_LIMBS - 1] = K_TOP_MIN + RAND64() % (P_TOP - K_TOP_MIN); + } + + // ── Step 2:设备端分配 ───────────────────────────────────────────────── + uint64_t *d_ct = nullptr, *d_pt = nullptr, *d_out = nullptr; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_pt, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_out, ct_bytes)); + + // ── Step 3:H2D(密文 + 明文),计时 ──────────────────────────────────── + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ciphers, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_pt, h_plains, exp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 4:GPU 密文乘明文 ──────────────────────────────────────────── + MulPlainTiming t = ou_mulplain( + d_ct, // 密文(蒙哥马利域) + d_pt, // 明文(22 × uint64_t/个) + d_out, // 输出(蒙哥马利域) + BATCH, d_r_0, d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ── Step 5:D2H 传输结果,计时 ───────────────────────────────────────── + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(h_results, d_out, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 6:打印计时 ─────────────────────────────────────────────────── + printf("\n========== ou_mulplain test ==========\n"); + printf(" batch=%d ARR_LEN=%d OU_TAU=%d OU_EXP_LIMBS=%d\n", BATCH, + ARR_LEN, OU_TAU, OU_EXP_LIMBS); + printf(" ct+pt H2D : %8.4f ms (%6.4f ms/op)\n", ms_h2d, ms_h2d / BATCH); + printf(" GPU kernel : %8.4f ms (%6.4f ms/op) [最长流]\n", t.kernel_ms, + t.kernel_ms / BATCH); + printf(" pipeline : %8.4f ms (%6.4f ms/op) [端到端]\n", t.pipeline_ms, + t.pipeline_ms / BATCH); + printf(" result D2H : %8.4f ms (%6.4f ms/op)\n", ms_d2h, ms_d2h / BATCH); + printf(" total/op : %6.4f ms\n", + (ms_h2d + t.pipeline_ms + ms_d2h) / BATCH); + printf("======================================\n\n"); + + /* + // ── Step 7:写入测试数据文件 ──────────────────────────────────────────── + // 格式(ou_mulplain_test.txt): + // 行 1(注释):元信息 + // 行 2 (N:) :OU 模数 n,256 limb,base-2^17 + // 行 3 (P:) :私钥 p,22 limb,base-2^64 + // 此后每用例 3 行:CT / PT / R + FILE *fp = fopen("ou_mulplain_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_mulplain_test.txt\n"); + } else { + fprintf(fp, + "# ou_mulplain test BATCH=%d ARR_LEN=%d BASE_BITS=%d" + " OU_EXP_LIMBS=%d CT_BASE=2^17 PT_BASE=2^64\n", + BATCH, ARR_LEN, BASE_BITS, OU_EXP_LIMBS); + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + fprintf(fp, "P:"); + for (int j = 0; j < OU_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)p_limbs[j]); + fprintf(fp, "\n"); + for (int b = 0; b < BATCH; b++) { + const uint64_t *ct = h_ciphers + (size_t)b * ARR_LEN; + const uint64_t *pt = h_plains + (size_t)b * OU_EXP_LIMBS; + const uint64_t *r = h_results + (size_t)b * ARR_LEN; + fprintf(fp, "CT:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ct[j]); + fprintf(fp, "\n"); + fprintf(fp, "PT:"); + for (int j = 0; j < OU_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)pt[j]); + fprintf(fp, "\n"); + fprintf(fp, "R:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)r[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_mulplain test] 测试数据已写入 ou_mulplain_test.txt\n"); + } + */ + + // ── 释放 ─────────────────────────────────────────────────────────────── +#undef RAND64 + free(h_ciphers); + free(h_plains); + free(h_results); + cudaFree(d_ct); + cudaFree(d_pt); + cudaFree(d_out); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_randc.cu b/heu/library/algorithms/ou_new/ou_randc.cu new file mode 100644 index 0000000..b029c70 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_randc.cu @@ -0,0 +1,19409 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 11687101653036, 18446743758138104250, 14211471509222940197, + 4235272491137479070, 18240298465959804861, 3612012136679214338, + 17250142769850486366, 16237778209286303069, 4059562170797630074, + 9869391574692501418, 11541833079389403685, 6455235399891900791, + 12385636973677184651, 16642176613718519917, 18371701685023660666, + 12908184103137142305, 432025558415355680, 13421505415041197067, + 3981525027614746084, 11947396656646251748, 17803480186391939853, + 8964175857302805981, 10755671907831908778, 1904820253353466484, + 10527642910445979147, 13113465356235953778, 7234975774234226545, + 9757071334591424272, 7316737729072523157, 6883729646458586186, + 8932421558757596044, 14597307983516136610, 14233691893978885667, + 11827243429966767869, 4500161214914298038, 10200789522299270510, + 13311736198424319672, 15003914107128413955, 15970198162388647651, + 5178000660069144055, 12257227048244175003, 199311466455739912, + 1199638074941611369, 7733792994982443122, 5589666049742506788, + 13804186403958915190, 17844357954872068228, 1608648031291287388, + 14833453796159433687, 17508515457688533059, 5757945642895465137, + 13081582882382389324, 5394527028006066918, 276092650195297071, + 17347510335080686092, 1701269284563161833, 13303804442297711418, + 8253121806998843455, 1803714610749342533, 15051344875346831329, + 17235261944528818002, 8632334691816709691, 7259437303239191782, + 15692170915486480673, 16909097158836466193, 3813100579756643340, + 8120672335331207645, 17658082942068012540, 14712625527555672008, + 3013490167685507391, 17053781224112072993, 7951833678156564288, + 16882129623470333590, 16598623833219974769, 8844055318436626140, + 7109452390029814739, 17973202994978822509, 17988667507327805934, + 4123463954687411915, 4940157814071223716, 10034667236662981698, + 2823422309362454166, 8801094461519448259, 16352417325266577247, + 4120137837347666316, 13914619055383639681, 13924524887541743843, + 16792129803426718801, 17665685411675475587, 4424050859418317508, + 7334727500762095896, 17503735803347907081, 17481354675564559188, + 3960789173235012097, 2382836270529746228, 4344431928835757414, + 10028103253977297380, 10742172769839739841, 4046021704283408275, + 15048763124930553406, 12884462872804057858, 6406919215710243952, + 2305936473049040111, 13121402678735537817, 4908133312471037493, + 11377681924055600462, 12660562847764485932, 4317962794267615766, + 1681209879049333151, 17555975506207345043, 5125753427319466102, + 3335447880892075219, 8356915887911857374, 8584812879987417990, + 17895049134452250854, 9085212301170466888, 7267866673654176139, + 11010976396069557933, 3608178248288276404, 676841772753514930, + 14867830803115014463, 1834280874555657868, 16310358636794094835, + 7330665673582977596, 15741791681143414831, 210676699798228406, + 17198551982049102727, 6109711417879930926, 10546103406410231004, + 15078747006884146927, 15249364729241398859, 7659463200845052268, + 6442795927427660859, 2250605931405808055, 8092318475578226474, + 18259756362431830946, 17518863421902388405, 10430337473484617894, + 11467857499548239759, 14850024957617392139, 4520997378651548347, + 1002071158803202119, 13705130616563662128, 16248739479905188442, + 3181222670542328118, 10465632553806001024, 13994891411389346011, + 10984874395183979750, 9226503659276477277, 16804196055515371214, + 10159135197864231902, 2843438330111370559, 10398977526462394360, + 13629959477406403811, 17564676539751491213, 9169917922917002660, + 9102739415085474595, 2571556159645270064, 15524480249380336561, + 16515752187799014418, 15631770625703612855, 15987278742587054296, + 6287084574542908707, 5110496278831165043, 3153708236541173992, + 12927769102029794814, 15247894309294568262, 9307521059673973398, + 2770646931232435574, 16454025378125412919, 10633343977751814089, + 1784373777963240394, 6934373261888962825, 17349218505508302052, + 10968015619489026254, 1752595997967202555, 12837659792768865327, + 14742784040235182007, 10962127976630829136, 5900158086628228693, + 17316277940499271723, 348494413055888541, 488358709608098659, + 10382164829144562111, 10796492275262689178, 10517686393808888673, + 779545184377220173, 3381663347793212032, 15236282164919686367, + 6334307549334192514, 3063003522052060686, 12114810039050492787, + 8870556826759400006, 4038453701720268007, 14314379071943608840, + 4339980657355467038, 9171890896160995321, 7917821014549284449, + 13180571956635383946, 18248798750186441796, 14577763713235528379, + 1809799949882400184, 6318379507872769143, 1709904138639204836, + 4433595655137503557, 14198791200961044293, 4650959752702308354, + 3318171583106947229, 209694021303727738, 1076839995989163599, + 16905707527716642672, 3695319746732430913, 10252674333030518371, + 16420700828256874086, 2433845634309674122, 16595131843996099370, + 16829576163922136336, 4841410023332049560, 12592434475652961886, + 17572224200405592702, 2431385938003789647, 10061028979483934059, + 2925075822122586101, 8606434114160323337, 5607119374066730577, + 11884780541782053893, 3126131661631420656, 10027052590555524485, + 16048091853305621015, 15852396680435263215, 17385109108245871297, + 12899005442699559936, 7848549015331456901, 9096729734807002481, + 4996467004929486051, 12243245730936161727, 9057745574396783344, + 7771655314147204603, 17881871823990369590, 15212325419875966733, + 11042754214829301385, 3281380445998958437, 17088850123667971831, + 14495632125498788619, 17994272834268936450, 16829150404837372, + 6402480299793703941, 10533393325012975763, 2416806924432625423, + 2875142845742952402, 12269921477737603466, 712826757316093295, + 4415075740707273176, 7975119161839365838, 11673666813117204770, + 9840315584201792241}; + +const uint64_t con_modn_shoup[256] = { + 14252880640204352951, 18322338132012530560, 6849211572557864124, + 15932334426407894751, 14177538223384229049, 9106170071204927292, + 16758578487033067404, 5864117390783715525, 16494545899140440021, + 5897258617717822505, 1933174352700629569, 12810083258791009448, + 12690514865985841899, 3970720354745169798, 4239814533413767079, + 5609102486863112468, 11230284723594595426, 17034417004294615591, + 5986132948557996967, 1868566874544028188, 2158239585541928173, + 17841097290863850509, 9305374060203424222, 10083694270531949160, + 2654649954734757684, 17721823101598855865, 599980504891318557, + 17732121566266018424, 1832524248816892725, 5295674783104160211, + 13283213776815003617, 10900691717351424196, 6021057974446928650, + 12624795618036119368, 2798162278377969124, 7399862538297940480, + 576839721297220460, 16060704215571397670, 16380205270440154947, + 7499979448419237887, 13254841413012858481, 10664669973443882596, + 1765312882701300737, 9426266066319221551, 12823009007753704482, + 10630699336822577567, 16298120910453621338, 13674950148586695572, + 17678273253225120972, 6798806775207395115, 13410427498759750653, + 2614783784077562964, 9342414901595102647, 14373786281595714575, + 6330183866169305354, 998675938268783033, 10221732541776071598, + 8979762078911881033, 9878667596621300249, 4856285279936479033, + 14833025980849776526, 5604878655902262946, 13088552648421720803, + 1801154013367199521, 4158119999782205988, 7904891504652660262, + 9042945763108429841, 4642264688478488779, 16204979912313018920, + 3580705517878336362, 9712754433060271621, 8675179099278674516, + 8186655178929728093, 1884659203003867161, 17775229374523385263, + 6390348527000753038, 6439058351892770174, 2339745637453507323, + 16274314407512660647, 2247518004490005028, 18003796786185432156, + 4540940947376355923, 11987538437574474975, 16166798012420960901, + 16121611900792272328, 16082928115738878740, 121528093926685229, + 11609994994605905995, 2593441955413327993, 16920803883743198476, + 11945409668615507125, 15459882499135165139, 4709903422099897132, + 1915945056478813527, 17487099108173624447, 4121351438621439846, + 11648490996845515622, 16906896413860707859, 440932069689474224, + 5596373384320545758, 6286719224840257488, 15070666469307485122, + 1718056780076659255, 16292491121877970301, 16399246121914763003, + 868264559834958645, 5880650461523548368, 13037697811177873232, + 12598280349103069353, 8787439026840426841, 75102682845531848, + 10793543124523682506, 4058772666671965704, 4575391113880810276, + 7977675084418792789, 3637051392050280908, 16362683407568863478, + 18347383388798481077, 9115743514592553391, 9421569851894468249, + 15594101773322529942, 11807267208082355523, 2211845703086696074, + 17348335706114235958, 11847926503719721254, 17547278040398801999, + 5869056178242350580, 13320003773654588467, 12699478824066277151, + 5239882100070320334, 7261256595809982529, 18328110662323898796, + 14262563528151153433, 10570694000294503318, 12813828885908106934, + 10929763919809758798, 2938308820079533814, 12010181661893483546, + 5724066348617601097, 14693589767406584560, 7346156359909105313, + 12463844683080763521, 6157213913132141689, 10056871538135474507, + 10533920527198962255, 8235152268042627670, 12319030087627737531, + 8540756947882872704, 7431325835456426550, 1301653294393655025, + 5378735902526386334, 14912383613060771114, 639721130790109100, + 16570337161183830109, 3674985562098081097, 25882515425358888, + 5781372063417524209, 16445334884700810773, 3544553957273777187, + 7670642182104980993, 9626872654279745485, 16105190074590295359, + 7770490841006776992, 15228876210409060859, 9849662374993193055, + 8654391918266496929, 6489400423788825431, 3633531112627186925, + 14858949636521671041, 1105232854426717343, 16217154593325743207, + 6106199821319202195, 15821125396653981131, 397434568115144159, + 5468761408652955936, 1296217405573136392, 8004677824854586354, + 2227606275875858042, 810603102045699119, 10604814613007946604, + 5290458938805336352, 17851909192937068431, 13268718299834334195, + 10806219279687765004, 7326952643401977865, 5984244256621982617, + 1659224770285885321, 10142490388121661620, 543184966312697096, + 6161334213393132777, 8606758596526178508, 16789552215120228336, + 8023727110355433065, 4647377966945102884, 6753109868947696463, + 3601294586865406454, 17607120286020078979, 2828322973879047780, + 2912791741152772408, 12563760677479960004, 14534280132691080513, + 1458078075305867271, 965960906590360464, 7718560025107401880, + 5982496231126478023, 9871084626240060187, 13176440103626612310, + 12254705932735020033, 6091017959275360856, 8575903195682037371, + 12248661790925770989, 15874428453561837902, 11580211822667122432, + 1675684581791044528, 953808119473653258, 5212502010992923248, + 11653707338854681558, 2814312282799796103, 10741977896772006037, + 792456711582471131, 5829394638712031862, 14582582386261703791, + 15116195068952376181, 16192690152961596921, 14996186982344279167, + 12452747948715197047, 9822110408961686525, 13084672463213903010, + 15511412028972873361, 4378034570765898628, 7434337709561426193, + 10855303285220731112, 7759227917166922418, 4976939851858003292, + 5204453497818198107, 16768838491807377833, 9561958848337334713, + 703798444570847189, 14816217796224051499, 2968875028840941609, + 6519175664834707328, 17450997194073458375, 12811758118865484166, + 10759621678827047865, 11859099579116844022, 14425180740705110568, + 7511257586748720299, 3539736587530686592, 4447312216206097890, + 11184913710261542236, 8771977661991653642, 12354338902272316516, + 15541265547581959679, 14587017360825887711, 15248329331280087903, + 10992628732632760992}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +// 保留参数兼容调用处,128-bit 时不需要 p +static void ou_gen_r_prime(const uint64_t * /*ou_p_limbs17*/, int batch, + uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 43351, 84159, 17007, 126963, 115814, 64975, 14865, 122878, 58093, + 76773, 16638, 87086, 110462, 105466, 35053, 36095, 8051, 116177, + 119699, 118157, 12357, 71314, 68424, 35266, 58013, 63468, 22117, + 10903, 124058, 90359, 68490, 117774, 56449, 45990, 26837, 86153, + 120741, 31603, 78596, 24019, 45134, 33649, 61458, 59406, 88868, + 60745, 113313, 123484, 30017, 98185, 93108, 73040, 39521, 18181, + 2647, 51647, 10194, 73702, 22934, 64, 29664, 94536, 9414, + 63827, 6028, 107137, 71399, 49216, 8196, 46100, 117329, 67195, + 25041, 122567, 110161, 82524, 85064, 85420, 38367, 90728, 6216, + 87366, 124652, 29067, 100922, 38894, 64688, 22860, 83774, 130371, + 39036, 94816, 45277, 76221, 67984, 78245, 70889, 64430, 52640, + 50933, 54580, 32496, 95587, 110988, 102834, 68631, 42744, 111149, + 127114, 116295, 108662, 4710, 31837, 15424, 50234, 99229, 61393, + 81585, 33195, 14128, 9168, 55047, 119038, 97329, 43164, 111637, + 39396, 13009, 90209, 92184, 81272, 101938, 57149, 82121, 100630, + 37780, 7881, 13181, 8505, 125111, 43862, 119168, 19431, 80034, + 114187, 71294, 52911, 81495, 14533, 87246, 126978, 30310, 9978, + 44551, 60081, 126942, 75376, 77030, 36034, 104993, 58885, 90371, + 111023, 45378, 97203, 126393, 72942, 8192, 124336, 37338, 116797, + 66693, 60337, 12040, 90738, 108119, 66171, 78981, 79494, 91989, + 89494, 118041, 29798, 30883, 110522, 122729, 7823, 62523, 20666, + 52089, 43045, 51146, 24317, 38753, 122735, 100047, 56716, 117911, + 60032, 27220, 44093, 56, 113553, 49629, 84418, 64845, 67097, + 10050, 8296, 90055, 75973, 63190, 82919, 56713, 30800, 90227, + 63208, 39501, 61899, 129744, 78395, 58460, 121961, 72489, 20054, + 64673, 102069, 663, 109348, 36701, 18676, 98666, 77634, 108307, + 103092, 49932, 49372, 39729, 72337, 90146, 247, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49877, 103520, 14633, 105196, 88121, 75472, 97344, 77893, 104350, + 114542, 91892, 50952, 64070, 21857, 78639, 2966, 6840, 35973, + 24928, 25085, 49586, 41185, 13904, 75536, 2710, 53110, 104559, + 109441, 40809, 28751, 119672, 47293, 31768, 78934, 55525, 52600, + 69445, 27490, 96611, 84604, 11133, 106316, 35838, 34416, 127179, + 34651, 118161, 16314, 45331, 107582, 24235, 12592, 81941, 19278, + 85161, 104589, 13239, 111883, 81763, 64318, 45287, 80867, 37986, + 126633, 118333, 32371, 49588, 111307, 6225, 98536, 50417, 128492, + 7707, 81120, 51703, 13989, 14502, 86184, 74055, 95503, 80070, + 61934, 73173, 322, 128915, 100622, 92460, 3170, 90102, 44305, + 79192, 110628, 84896, 92155, 13698, 129632, 82311, 123141, 99527, + 28216, 52443, 78019, 37062, 11803, 15622, 2177, 42554, 81945, + 37634, 97471, 11261, 30170, 62893, 129991, 77778, 123677, 75667, + 22518, 67693, 79122, 34634, 117067, 33393, 41664, 46880, 63508, + 1202, 120484, 35970, 14884, 64323, 87199, 61041, 17700, 6496, + 128561, 104072, 129267, 31935, 119451, 54205, 55539, 59285, 49953, + 34036, 45769, 28494, 30664, 79616, 47756, 57833, 99317, 122074, + 130764, 65477, 42562, 43963, 97923, 114657, 34735, 19193, 84201, + 24821, 98149, 101368, 100860, 110862, 96240, 101375, 55675, 99994, + 12323, 56026, 120364, 84030, 10327, 108568, 4795, 122128, 3775, + 1740, 113334, 5740, 61052, 15255, 44939, 84950, 4631, 87197, + 63464, 45041, 49844, 102052, 41710, 76594, 96748, 2213, 15033, + 56862, 42121, 22702, 29104, 55955, 32193, 62378, 61812, 37549, + 27929, 118796, 116386, 35884, 83278, 116744, 103768, 106752, 29801, + 6976, 81713, 55669, 12038, 51733, 6915, 128541, 82038, 102167, + 64630, 125581, 69829, 79662, 80895, 89416, 41571, 113918, 73736, + 22655, 72892, 97009, 75512, 83469, 50798, 35893, 72631, 27752, + 114176, 116066, 35199, 11556, 117400, 53979, 71662, 76589, 35790, + 51797, 38295, 48839, 44050}; + uint64_t R1[256] = { + 95787, 87855, 74590, 120089, 96462, 125333, 89873, 62820, 56744, + 93675, 114260, 58407, 55044, 4742, 20922, 129032, 18634, 103071, + 2852, 114517, 116272, 79216, 95365, 35177, 128432, 80425, 18923, + 592, 13588, 42144, 48019, 39668, 66805, 42663, 33194, 65911, + 93428, 9610, 76041, 48300, 121686, 67062, 30099, 86626, 99273, + 90908, 66468, 28073, 71719, 43868, 40340, 88274, 109318, 21824, + 16472, 116161, 1346, 106033, 20342, 35258, 20632, 105594, 118266, + 97653, 97643, 7306, 67863, 79950, 112151, 117205, 39906, 100559, + 19328, 14826, 43881, 64539, 123341, 37113, 15909, 43631, 27755, + 12868, 87791, 110907, 2763, 41576, 76238, 21079, 28709, 67173, + 22692, 45867, 80137, 111080, 90017, 19215, 15056, 7843, 34411, + 10397, 47157, 113197, 77959, 43337, 123310, 26898, 103324, 95568, + 37773, 53742, 58532, 64900, 12429, 109482, 75505, 70429, 89935, + 67404, 103144, 45403, 28839, 100826, 80183, 60279, 60825, 67114, + 15456, 95163, 5820, 106812, 38605, 127798, 43023, 23037, 109334, + 82354, 36764, 29882, 1460, 109709, 70002, 61938, 129339, 12574, + 23578, 116415, 67219, 51854, 56951, 3482, 98043, 101818, 116934, + 91679, 101444, 12152, 24202, 116763, 23200, 65030, 72892, 11401, + 89306, 104758, 94473, 69024, 2331, 120575, 38998, 71116, 123520, + 12246, 31511, 15417, 54824, 35449, 51579, 129451, 119392, 117118, + 124644, 45532, 83498, 101978, 4264, 7165, 33903, 43033, 21694, + 89206, 4834, 104846, 13789, 111463, 68368, 35303, 42736, 93343, + 26110, 67668, 100827, 18175, 80987, 95026, 43049, 29875, 73340, + 103134, 127164, 12263, 68500, 44986, 40419, 118430, 83125, 27361, + 83899, 599, 102652, 96223, 98631, 95859, 60979, 40494, 118153, + 71598, 113945, 125768, 2378, 115813, 13443, 7038, 3598, 53377, + 26625, 110573, 3294, 119847, 72218, 85832, 125, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 2023, 29313, 55760, 30415, 117209, 68859, 126088, 105595, 44062, + 129614, 108803, 62812, 119884, 130526, 4554, 108683, 53422, 114202, + 34719, 66446, 65067, 23670, 9591, 82680, 20040, 77609, 34903, + 85401, 37760, 43899, 14395, 67595, 12710, 101806, 77200, 103316, + 65099, 125953, 105217, 67808, 5979, 63677, 12355, 111674, 70131, + 45594, 117232, 10425, 81356, 112691, 14758, 10490, 24556, 27922, + 32980, 20928, 118420, 17204, 4244, 126937, 116849, 106497, 51321, + 114935, 45503, 15461, 59271, 111583, 30113, 103352, 10622, 32510, + 41116, 86928, 10137, 101567, 30707, 124356, 108755, 8203, 11158, + 9603, 114740, 5093, 13054, 61800, 75687, 38080, 11550, 87153, + 33247, 18929, 66437, 13511, 39575, 18765, 61155, 77315, 112366, + 76906, 125693, 40793, 40582, 43161, 81338, 111531, 84813, 49322, + 71309, 83250, 14948, 44745, 13967, 98243, 116072, 5842, 82567, + 77993, 80649, 107659, 66320, 122438, 54394, 54983, 79006, 105681, + 94840, 79085, 41950, 106863, 130420, 89427, 83726, 86511, 44750, + 12837, 47751, 33678, 115313, 66053, 43941, 100068, 107956, 169, + 62673, 70167, 106884, 28460, 37125, 129538, 93387, 56010, 35558, + 40621, 77463, 39765, 16013, 85203, 39465, 19946, 122865, 58068, + 60861, 54036, 48234, 130529, 59321, 83170, 116672, 11733, 94357, + 35207, 46800, 22992, 47306, 80973, 28136, 59828, 4338, 109019, + 16604, 58136, 123247, 75151, 11948, 105570, 86992, 72951, 15873, + 12088, 27387, 60538, 80333, 72819, 97245, 41421, 126078, 82127, + 78747, 70199, 99633, 38926, 56088, 20821, 130509, 125699, 127559, + 47166, 102354, 71946, 79468, 81253, 11120, 37016, 20384, 95174, + 126298, 33756, 60551, 106592, 43758, 59029, 22997, 63753, 65367, + 47855, 5100, 118402, 102326, 108278, 49481, 99388, 114710, 101475, + 124482, 27406, 78309, 52397, 69291, 787, 10, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ========================================================================= + // ou_randomize 正确性测试 + // + // ou_randomize(c̃, r') = c̃ · H^{r'} mod n(FMLM 域,原地) + // + // 流程: + // 1. 建立 H 预计算表(CPU → GPU) + // 2. 随机生成 BATCH 个密文 c̃(< n,FMLM 域) + // 3. CPU 生成随机指数 r'(ou_gen_r_prime,r' < p) + // 4. H2D:c̃ + r' + // 5. GPU:ou_randomize → d_c 原地更新 + // 6. D2H:结果 + // 7. 写入 ou_randomize_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 200000; + + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const size_t rp_bytes = (size_t)BATCH * OU_HR_EXP_LIMBS * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + const uint64_t MASK17 = (1ULL << BASE_BITS) - 1ULL; + const uint64_t CT_TOP = Modn[240]; // = 247,密文最高 limb 严格上界 + + srand((unsigned int)time(nullptr)); + + // ── Step 1:建立 H 预计算表(CPU),上传 GPU ──────────────────────────── + uint64_t *h_H_table = (uint64_t *)malloc(tbl_bytes); + generate_H_table(Modn, ou_H, h_H_table); + + uint64_t *d_H_table = nullptr; + CUDA_CHECK(cudaMalloc(&d_H_table, tbl_bytes)); + CUDA_CHECK( + cudaMemcpy(d_H_table, h_H_table, tbl_bytes, cudaMemcpyHostToDevice)); + free(h_H_table); + + // ── Step 2:随机生成 BATCH 个 c̃(视作 FMLM 域密文,满足 0 < c̃ < n) ── + uint64_t *h_c_orig = (uint64_t *)malloc(ct_bytes); + for (int b = 0; b < BATCH; b++) { + uint64_t *c = h_c_orig + (size_t)b * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < 240) { + c[j] = (((uint64_t)(uint32_t)rand()) ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17; + } else if (j == 240) { + c[j] = (uint64_t)rand() % CT_TOP; + } else { + c[j] = 0ULL; + } + } + } + + // ── Step 3:CPU 生成随机指数 r'(OU_HR_TAU=128 + // bit,base-2^64,OU_HR_EXP_LIMBS=2)── + uint64_t *h_rp = (uint64_t *)malloc(rp_bytes); + ou_gen_r_prime(ou_p, BATCH, h_rp); + + // ── Step 4:H2D,计时 ──────────────────────────────────────────────────── + uint64_t *d_c = nullptr, *d_rp = nullptr; + CUDA_CHECK(cudaMalloc(&d_c, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_rp, rp_bytes)); + + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(d_c, h_c_orig, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_rp, h_rp, rp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 5:GPU ou_randomize,计时 ─────────────────────────────────────── + float ms_gpu = ou_randomize( + d_c, d_H_table, d_rp, BATCH, d_negModn, d_con_NegModn_shoup, d_Modn, + d_con_Modn_shoup, MOD, d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, + inv_shoup, d_sample); + + // ── Step 6:D2H,计时 ──────────────────────────────────────────────────── + uint64_t *h_c_new = (uint64_t *)malloc(ct_bytes); + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(h_c_new, d_c, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 7:打印计时 + // ────────────────────────────────────────────────────── + printf("\n========== ou_randomize test ==========\n"); + printf(" batch=%d ARR_LEN=%d BASE_BITS=%d OU_HR_TAU=%d\n", BATCH, + ARR_LEN, BASE_BITS, OU_HR_TAU); + printf(" [H2D c+r'] : %8.4f ms (%6.4f ms/op)\n", ms_h2d, + ms_h2d / BATCH); + printf(" [GPU ou_randomize] : %8.4f ms (%6.4f ms/op)\n", ms_gpu, + ms_gpu / BATCH); + printf(" [D2H c_new] : %8.4f ms (%6.4f ms/op)\n", ms_d2h, + ms_d2h / BATCH); + printf("=======================================\n\n"); + /* + // ── Step 8:写测试数据文件 + ──────────────────────────────────────────────── + // 格式 ou_randomize_test.txt: + // 第 1 行 (注释):元信息 + // 第 2 行 N: OU 模数 n(256 limb,base-2^17) + // 第 3 行 H: 公钥 H (256 limb,base-2^17) + // 每用例 3 行: + // C: 原始密文 c̃(256 limb,base-2^17,FMLM 域) + // RP: 随机指数 r'(OU_HR_EXP_LIMBS=2 + limb,base-2^64,小端序,OU_HR_TAU=128 bit) + // CN: 随机化后密文 c̃_new(256 limb,base-2^17,FMLM 域) + FILE *fp = fopen("ou_randomize_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_randomize_test.txt\n"); + } else { + fprintf(fp, + "# ou_randomize test BATCH=%d ARR_LEN=%d BASE_BITS=%d" + " OU_HR_TAU=%d OU_HR_EXP_LIMBS=%d\n", + BATCH, ARR_LEN, BASE_BITS, OU_HR_TAU, OU_HR_EXP_LIMBS); + // N: modulus n(base-2^17,256 limb) + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + // H: public key H(base-2^17,256 limb) + fprintf(fp, "H:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)ou_H[j]); + fprintf(fp, "\n"); + // 逐用例写 C / RP / CN + for (int b = 0; b < BATCH; b++) { + const uint64_t *c_orig = h_c_orig + (size_t)b * ARR_LEN; + const uint64_t *rp = h_rp + (size_t)b * OU_HR_EXP_LIMBS; + const uint64_t *c_new = h_c_new + (size_t)b * ARR_LEN; + // C: 原始密文(base-2^17,256 limb) + fprintf(fp, "C:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c_orig[j]); + fprintf(fp, "\n"); + // RP: 随机指数 r'(base-2^64,OU_HR_EXP_LIMBS=2 limb) + fprintf(fp, "RP:"); + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + fprintf(fp, " %llu", (unsigned long long)rp[j]); + fprintf(fp, "\n"); + // CN: 随机化后密文(base-2^17,256 limb) + fprintf(fp, "CN:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c_new[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_randomize test] 测试数据已写入 ou_randomize_test.txt\n"); + } + */ + // ── 释放 + // ────────────────────────────────────────────────────────────────── + free(h_c_orig); + free(h_rp); + free(h_c_new); + cudaFree(d_c); + cudaFree(d_rp); + cudaFree(d_H_table); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/ou_new/ou_subhomo.cu b/heu/library/algorithms/ou_new/ou_subhomo.cu new file mode 100644 index 0000000..402e2a8 --- /dev/null +++ b/heu/library/algorithms/ou_new/ou_subhomo.cu @@ -0,0 +1,19424 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 11687101653036, 18446743758138104250, 14211471509222940197, + 4235272491137479070, 18240298465959804861, 3612012136679214338, + 17250142769850486366, 16237778209286303069, 4059562170797630074, + 9869391574692501418, 11541833079389403685, 6455235399891900791, + 12385636973677184651, 16642176613718519917, 18371701685023660666, + 12908184103137142305, 432025558415355680, 13421505415041197067, + 3981525027614746084, 11947396656646251748, 17803480186391939853, + 8964175857302805981, 10755671907831908778, 1904820253353466484, + 10527642910445979147, 13113465356235953778, 7234975774234226545, + 9757071334591424272, 7316737729072523157, 6883729646458586186, + 8932421558757596044, 14597307983516136610, 14233691893978885667, + 11827243429966767869, 4500161214914298038, 10200789522299270510, + 13311736198424319672, 15003914107128413955, 15970198162388647651, + 5178000660069144055, 12257227048244175003, 199311466455739912, + 1199638074941611369, 7733792994982443122, 5589666049742506788, + 13804186403958915190, 17844357954872068228, 1608648031291287388, + 14833453796159433687, 17508515457688533059, 5757945642895465137, + 13081582882382389324, 5394527028006066918, 276092650195297071, + 17347510335080686092, 1701269284563161833, 13303804442297711418, + 8253121806998843455, 1803714610749342533, 15051344875346831329, + 17235261944528818002, 8632334691816709691, 7259437303239191782, + 15692170915486480673, 16909097158836466193, 3813100579756643340, + 8120672335331207645, 17658082942068012540, 14712625527555672008, + 3013490167685507391, 17053781224112072993, 7951833678156564288, + 16882129623470333590, 16598623833219974769, 8844055318436626140, + 7109452390029814739, 17973202994978822509, 17988667507327805934, + 4123463954687411915, 4940157814071223716, 10034667236662981698, + 2823422309362454166, 8801094461519448259, 16352417325266577247, + 4120137837347666316, 13914619055383639681, 13924524887541743843, + 16792129803426718801, 17665685411675475587, 4424050859418317508, + 7334727500762095896, 17503735803347907081, 17481354675564559188, + 3960789173235012097, 2382836270529746228, 4344431928835757414, + 10028103253977297380, 10742172769839739841, 4046021704283408275, + 15048763124930553406, 12884462872804057858, 6406919215710243952, + 2305936473049040111, 13121402678735537817, 4908133312471037493, + 11377681924055600462, 12660562847764485932, 4317962794267615766, + 1681209879049333151, 17555975506207345043, 5125753427319466102, + 3335447880892075219, 8356915887911857374, 8584812879987417990, + 17895049134452250854, 9085212301170466888, 7267866673654176139, + 11010976396069557933, 3608178248288276404, 676841772753514930, + 14867830803115014463, 1834280874555657868, 16310358636794094835, + 7330665673582977596, 15741791681143414831, 210676699798228406, + 17198551982049102727, 6109711417879930926, 10546103406410231004, + 15078747006884146927, 15249364729241398859, 7659463200845052268, + 6442795927427660859, 2250605931405808055, 8092318475578226474, + 18259756362431830946, 17518863421902388405, 10430337473484617894, + 11467857499548239759, 14850024957617392139, 4520997378651548347, + 1002071158803202119, 13705130616563662128, 16248739479905188442, + 3181222670542328118, 10465632553806001024, 13994891411389346011, + 10984874395183979750, 9226503659276477277, 16804196055515371214, + 10159135197864231902, 2843438330111370559, 10398977526462394360, + 13629959477406403811, 17564676539751491213, 9169917922917002660, + 9102739415085474595, 2571556159645270064, 15524480249380336561, + 16515752187799014418, 15631770625703612855, 15987278742587054296, + 6287084574542908707, 5110496278831165043, 3153708236541173992, + 12927769102029794814, 15247894309294568262, 9307521059673973398, + 2770646931232435574, 16454025378125412919, 10633343977751814089, + 1784373777963240394, 6934373261888962825, 17349218505508302052, + 10968015619489026254, 1752595997967202555, 12837659792768865327, + 14742784040235182007, 10962127976630829136, 5900158086628228693, + 17316277940499271723, 348494413055888541, 488358709608098659, + 10382164829144562111, 10796492275262689178, 10517686393808888673, + 779545184377220173, 3381663347793212032, 15236282164919686367, + 6334307549334192514, 3063003522052060686, 12114810039050492787, + 8870556826759400006, 4038453701720268007, 14314379071943608840, + 4339980657355467038, 9171890896160995321, 7917821014549284449, + 13180571956635383946, 18248798750186441796, 14577763713235528379, + 1809799949882400184, 6318379507872769143, 1709904138639204836, + 4433595655137503557, 14198791200961044293, 4650959752702308354, + 3318171583106947229, 209694021303727738, 1076839995989163599, + 16905707527716642672, 3695319746732430913, 10252674333030518371, + 16420700828256874086, 2433845634309674122, 16595131843996099370, + 16829576163922136336, 4841410023332049560, 12592434475652961886, + 17572224200405592702, 2431385938003789647, 10061028979483934059, + 2925075822122586101, 8606434114160323337, 5607119374066730577, + 11884780541782053893, 3126131661631420656, 10027052590555524485, + 16048091853305621015, 15852396680435263215, 17385109108245871297, + 12899005442699559936, 7848549015331456901, 9096729734807002481, + 4996467004929486051, 12243245730936161727, 9057745574396783344, + 7771655314147204603, 17881871823990369590, 15212325419875966733, + 11042754214829301385, 3281380445998958437, 17088850123667971831, + 14495632125498788619, 17994272834268936450, 16829150404837372, + 6402480299793703941, 10533393325012975763, 2416806924432625423, + 2875142845742952402, 12269921477737603466, 712826757316093295, + 4415075740707273176, 7975119161839365838, 11673666813117204770, + 9840315584201792241}; + +const uint64_t con_modn_shoup[256] = { + 14252880640204352951, 18322338132012530560, 6849211572557864124, + 15932334426407894751, 14177538223384229049, 9106170071204927292, + 16758578487033067404, 5864117390783715525, 16494545899140440021, + 5897258617717822505, 1933174352700629569, 12810083258791009448, + 12690514865985841899, 3970720354745169798, 4239814533413767079, + 5609102486863112468, 11230284723594595426, 17034417004294615591, + 5986132948557996967, 1868566874544028188, 2158239585541928173, + 17841097290863850509, 9305374060203424222, 10083694270531949160, + 2654649954734757684, 17721823101598855865, 599980504891318557, + 17732121566266018424, 1832524248816892725, 5295674783104160211, + 13283213776815003617, 10900691717351424196, 6021057974446928650, + 12624795618036119368, 2798162278377969124, 7399862538297940480, + 576839721297220460, 16060704215571397670, 16380205270440154947, + 7499979448419237887, 13254841413012858481, 10664669973443882596, + 1765312882701300737, 9426266066319221551, 12823009007753704482, + 10630699336822577567, 16298120910453621338, 13674950148586695572, + 17678273253225120972, 6798806775207395115, 13410427498759750653, + 2614783784077562964, 9342414901595102647, 14373786281595714575, + 6330183866169305354, 998675938268783033, 10221732541776071598, + 8979762078911881033, 9878667596621300249, 4856285279936479033, + 14833025980849776526, 5604878655902262946, 13088552648421720803, + 1801154013367199521, 4158119999782205988, 7904891504652660262, + 9042945763108429841, 4642264688478488779, 16204979912313018920, + 3580705517878336362, 9712754433060271621, 8675179099278674516, + 8186655178929728093, 1884659203003867161, 17775229374523385263, + 6390348527000753038, 6439058351892770174, 2339745637453507323, + 16274314407512660647, 2247518004490005028, 18003796786185432156, + 4540940947376355923, 11987538437574474975, 16166798012420960901, + 16121611900792272328, 16082928115738878740, 121528093926685229, + 11609994994605905995, 2593441955413327993, 16920803883743198476, + 11945409668615507125, 15459882499135165139, 4709903422099897132, + 1915945056478813527, 17487099108173624447, 4121351438621439846, + 11648490996845515622, 16906896413860707859, 440932069689474224, + 5596373384320545758, 6286719224840257488, 15070666469307485122, + 1718056780076659255, 16292491121877970301, 16399246121914763003, + 868264559834958645, 5880650461523548368, 13037697811177873232, + 12598280349103069353, 8787439026840426841, 75102682845531848, + 10793543124523682506, 4058772666671965704, 4575391113880810276, + 7977675084418792789, 3637051392050280908, 16362683407568863478, + 18347383388798481077, 9115743514592553391, 9421569851894468249, + 15594101773322529942, 11807267208082355523, 2211845703086696074, + 17348335706114235958, 11847926503719721254, 17547278040398801999, + 5869056178242350580, 13320003773654588467, 12699478824066277151, + 5239882100070320334, 7261256595809982529, 18328110662323898796, + 14262563528151153433, 10570694000294503318, 12813828885908106934, + 10929763919809758798, 2938308820079533814, 12010181661893483546, + 5724066348617601097, 14693589767406584560, 7346156359909105313, + 12463844683080763521, 6157213913132141689, 10056871538135474507, + 10533920527198962255, 8235152268042627670, 12319030087627737531, + 8540756947882872704, 7431325835456426550, 1301653294393655025, + 5378735902526386334, 14912383613060771114, 639721130790109100, + 16570337161183830109, 3674985562098081097, 25882515425358888, + 5781372063417524209, 16445334884700810773, 3544553957273777187, + 7670642182104980993, 9626872654279745485, 16105190074590295359, + 7770490841006776992, 15228876210409060859, 9849662374993193055, + 8654391918266496929, 6489400423788825431, 3633531112627186925, + 14858949636521671041, 1105232854426717343, 16217154593325743207, + 6106199821319202195, 15821125396653981131, 397434568115144159, + 5468761408652955936, 1296217405573136392, 8004677824854586354, + 2227606275875858042, 810603102045699119, 10604814613007946604, + 5290458938805336352, 17851909192937068431, 13268718299834334195, + 10806219279687765004, 7326952643401977865, 5984244256621982617, + 1659224770285885321, 10142490388121661620, 543184966312697096, + 6161334213393132777, 8606758596526178508, 16789552215120228336, + 8023727110355433065, 4647377966945102884, 6753109868947696463, + 3601294586865406454, 17607120286020078979, 2828322973879047780, + 2912791741152772408, 12563760677479960004, 14534280132691080513, + 1458078075305867271, 965960906590360464, 7718560025107401880, + 5982496231126478023, 9871084626240060187, 13176440103626612310, + 12254705932735020033, 6091017959275360856, 8575903195682037371, + 12248661790925770989, 15874428453561837902, 11580211822667122432, + 1675684581791044528, 953808119473653258, 5212502010992923248, + 11653707338854681558, 2814312282799796103, 10741977896772006037, + 792456711582471131, 5829394638712031862, 14582582386261703791, + 15116195068952376181, 16192690152961596921, 14996186982344279167, + 12452747948715197047, 9822110408961686525, 13084672463213903010, + 15511412028972873361, 4378034570765898628, 7434337709561426193, + 10855303285220731112, 7759227917166922418, 4976939851858003292, + 5204453497818198107, 16768838491807377833, 9561958848337334713, + 703798444570847189, 14816217796224051499, 2968875028840941609, + 6519175664834707328, 17450997194073458375, 12811758118865484166, + 10759621678827047865, 11859099579116844022, 14425180740705110568, + 7511257586748720299, 3539736587530686592, 4447312216206097890, + 11184913710261542236, 8771977661991653642, 12354338902272316516, + 15541265547581959679, 14587017360825887711, 15248329331280087903, + 10992628732632760992}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',满足 0 ≤ r'_i < p(OU 私钥 p 的阶 = H 的阶)。 +// +// 算法: +// 1. 将 ou_p(base-2^17,ARR_LEN=256 limb)转换为 base-2^64(OU_EXP_LIMBS=22 +// limb) +// 2. 动态计算 p 的有效 bit 数及顶端 limb 掩码 +// 3. 对每个 r':生成 OU_TAU-bit 随机数,截断到 p 的有效位宽, +// 拒绝采样直到 r' < p(期望 <2 次采样) +// +// 参数: +// ou_p_limbs17 [ARR_LEN] OU 私钥 p(base-2^17 小端序,只读) +// batch 生成数量 +// h_r_prime [batch × OU_EXP_LIMBS] 输出随机指数(base-2^64 小端序 +// uint64) +// ============================================================================= +static void ou_gen_r_prime(const uint64_t *ou_p_limbs17, int batch, + uint64_t *h_r_prime) { + // 1. 将 p 从 base-2^17(256 limb)转换为 base-2^64(OU_EXP_LIMBS=22 limb) + uint64_t p64[OU_EXP_LIMBS] = {}; + { + uint64_t bit_pos = 0; + for (int i = 0; i < ARR_LEN; i++) { + int j = (int)(bit_pos / 64); + int b = (int)(bit_pos % 64); + if (j >= OU_EXP_LIMBS) break; + p64[j] |= ou_p_limbs17[i] << b; + if (b + BASE_BITS > 64 && j + 1 < OU_EXP_LIMBS) + p64[j + 1] |= ou_p_limbs17[i] >> (64 - b); + bit_pos += BASE_BITS; + } + } + + // 2. 动态计算 p 的有效 bit 数及顶端 limb 掩码 + int top_limb = -1; + for (int j = OU_EXP_LIMBS - 1; j >= 0; j--) { + if (p64[j] != 0) { + top_limb = j; + break; + } + } + if (top_limb < 0) { + printf("[ou_gen_r_prime] ERROR: p == 0\n"); + return; + } + { + // 打印 p 的 bit 长度(仅首次调用) + int cnt = 0; + uint64_t v = p64[top_limb]; + while (v > 0) { + v >>= 1; + ++cnt; + } + printf("[ou_gen_r_prime] p bit length = %d (top_limb=%d, top_bits=%d)\n", + top_limb * 64 + cnt, top_limb, cnt); + } + int top_bits = 0; + { + uint64_t v = p64[top_limb]; + while (v > 0) { + v >>= 1; + ++top_bits; + } + } + const uint64_t top_mask = + (top_bits == 64) ? UINT64_MAX : ((1ULL << top_bits) - 1ULL); + + // 3. 逐条拒绝采样:生成 r'_i 均匀分布于 [0, p-1] + // lt_p(r) = true 当且仅当 r < p(小端 uint64 逐 limb 从高到低比较) + auto lt_p = [&](const uint64_t *r) -> bool { + for (int j = OU_EXP_LIMBS - 1; j >= 0; j--) { + if (r[j] < p64[j]) return true; + if (r[j] > p64[j]) return false; + } + return false; // r == p,拒绝 + }; + + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_EXP_LIMBS; + do { + // 用 rand()(15 bit)拼装 64 bit 随机数,适配 Windows RAND_MAX=32767 + for (int j = 0; j < OU_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + // 截断 top_limb 以上的位 + r[top_limb] &= top_mask; + for (int j = top_limb + 1; j < OU_EXP_LIMBS; j++) r[j] = 0; + } while (!lt_p(r)); // 拒绝 r' >= p + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_EXP_LIMBS] 随机指数 +// r'(base-2^64 小端序,ou_gen_r_prime 生成) batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 43351, 84159, 17007, 126963, 115814, 64975, 14865, 122878, 58093, + 76773, 16638, 87086, 110462, 105466, 35053, 36095, 8051, 116177, + 119699, 118157, 12357, 71314, 68424, 35266, 58013, 63468, 22117, + 10903, 124058, 90359, 68490, 117774, 56449, 45990, 26837, 86153, + 120741, 31603, 78596, 24019, 45134, 33649, 61458, 59406, 88868, + 60745, 113313, 123484, 30017, 98185, 93108, 73040, 39521, 18181, + 2647, 51647, 10194, 73702, 22934, 64, 29664, 94536, 9414, + 63827, 6028, 107137, 71399, 49216, 8196, 46100, 117329, 67195, + 25041, 122567, 110161, 82524, 85064, 85420, 38367, 90728, 6216, + 87366, 124652, 29067, 100922, 38894, 64688, 22860, 83774, 130371, + 39036, 94816, 45277, 76221, 67984, 78245, 70889, 64430, 52640, + 50933, 54580, 32496, 95587, 110988, 102834, 68631, 42744, 111149, + 127114, 116295, 108662, 4710, 31837, 15424, 50234, 99229, 61393, + 81585, 33195, 14128, 9168, 55047, 119038, 97329, 43164, 111637, + 39396, 13009, 90209, 92184, 81272, 101938, 57149, 82121, 100630, + 37780, 7881, 13181, 8505, 125111, 43862, 119168, 19431, 80034, + 114187, 71294, 52911, 81495, 14533, 87246, 126978, 30310, 9978, + 44551, 60081, 126942, 75376, 77030, 36034, 104993, 58885, 90371, + 111023, 45378, 97203, 126393, 72942, 8192, 124336, 37338, 116797, + 66693, 60337, 12040, 90738, 108119, 66171, 78981, 79494, 91989, + 89494, 118041, 29798, 30883, 110522, 122729, 7823, 62523, 20666, + 52089, 43045, 51146, 24317, 38753, 122735, 100047, 56716, 117911, + 60032, 27220, 44093, 56, 113553, 49629, 84418, 64845, 67097, + 10050, 8296, 90055, 75973, 63190, 82919, 56713, 30800, 90227, + 63208, 39501, 61899, 129744, 78395, 58460, 121961, 72489, 20054, + 64673, 102069, 663, 109348, 36701, 18676, 98666, 77634, 108307, + 103092, 49932, 49372, 39729, 72337, 90146, 247, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49877, 103520, 14633, 105196, 88121, 75472, 97344, 77893, 104350, + 114542, 91892, 50952, 64070, 21857, 78639, 2966, 6840, 35973, + 24928, 25085, 49586, 41185, 13904, 75536, 2710, 53110, 104559, + 109441, 40809, 28751, 119672, 47293, 31768, 78934, 55525, 52600, + 69445, 27490, 96611, 84604, 11133, 106316, 35838, 34416, 127179, + 34651, 118161, 16314, 45331, 107582, 24235, 12592, 81941, 19278, + 85161, 104589, 13239, 111883, 81763, 64318, 45287, 80867, 37986, + 126633, 118333, 32371, 49588, 111307, 6225, 98536, 50417, 128492, + 7707, 81120, 51703, 13989, 14502, 86184, 74055, 95503, 80070, + 61934, 73173, 322, 128915, 100622, 92460, 3170, 90102, 44305, + 79192, 110628, 84896, 92155, 13698, 129632, 82311, 123141, 99527, + 28216, 52443, 78019, 37062, 11803, 15622, 2177, 42554, 81945, + 37634, 97471, 11261, 30170, 62893, 129991, 77778, 123677, 75667, + 22518, 67693, 79122, 34634, 117067, 33393, 41664, 46880, 63508, + 1202, 120484, 35970, 14884, 64323, 87199, 61041, 17700, 6496, + 128561, 104072, 129267, 31935, 119451, 54205, 55539, 59285, 49953, + 34036, 45769, 28494, 30664, 79616, 47756, 57833, 99317, 122074, + 130764, 65477, 42562, 43963, 97923, 114657, 34735, 19193, 84201, + 24821, 98149, 101368, 100860, 110862, 96240, 101375, 55675, 99994, + 12323, 56026, 120364, 84030, 10327, 108568, 4795, 122128, 3775, + 1740, 113334, 5740, 61052, 15255, 44939, 84950, 4631, 87197, + 63464, 45041, 49844, 102052, 41710, 76594, 96748, 2213, 15033, + 56862, 42121, 22702, 29104, 55955, 32193, 62378, 61812, 37549, + 27929, 118796, 116386, 35884, 83278, 116744, 103768, 106752, 29801, + 6976, 81713, 55669, 12038, 51733, 6915, 128541, 82038, 102167, + 64630, 125581, 69829, 79662, 80895, 89416, 41571, 113918, 73736, + 22655, 72892, 97009, 75512, 83469, 50798, 35893, 72631, 27752, + 114176, 116066, 35199, 11556, 117400, 53979, 71662, 76589, 35790, + 51797, 38295, 48839, 44050}; + uint64_t R1[256] = { + 95787, 87855, 74590, 120089, 96462, 125333, 89873, 62820, 56744, + 93675, 114260, 58407, 55044, 4742, 20922, 129032, 18634, 103071, + 2852, 114517, 116272, 79216, 95365, 35177, 128432, 80425, 18923, + 592, 13588, 42144, 48019, 39668, 66805, 42663, 33194, 65911, + 93428, 9610, 76041, 48300, 121686, 67062, 30099, 86626, 99273, + 90908, 66468, 28073, 71719, 43868, 40340, 88274, 109318, 21824, + 16472, 116161, 1346, 106033, 20342, 35258, 20632, 105594, 118266, + 97653, 97643, 7306, 67863, 79950, 112151, 117205, 39906, 100559, + 19328, 14826, 43881, 64539, 123341, 37113, 15909, 43631, 27755, + 12868, 87791, 110907, 2763, 41576, 76238, 21079, 28709, 67173, + 22692, 45867, 80137, 111080, 90017, 19215, 15056, 7843, 34411, + 10397, 47157, 113197, 77959, 43337, 123310, 26898, 103324, 95568, + 37773, 53742, 58532, 64900, 12429, 109482, 75505, 70429, 89935, + 67404, 103144, 45403, 28839, 100826, 80183, 60279, 60825, 67114, + 15456, 95163, 5820, 106812, 38605, 127798, 43023, 23037, 109334, + 82354, 36764, 29882, 1460, 109709, 70002, 61938, 129339, 12574, + 23578, 116415, 67219, 51854, 56951, 3482, 98043, 101818, 116934, + 91679, 101444, 12152, 24202, 116763, 23200, 65030, 72892, 11401, + 89306, 104758, 94473, 69024, 2331, 120575, 38998, 71116, 123520, + 12246, 31511, 15417, 54824, 35449, 51579, 129451, 119392, 117118, + 124644, 45532, 83498, 101978, 4264, 7165, 33903, 43033, 21694, + 89206, 4834, 104846, 13789, 111463, 68368, 35303, 42736, 93343, + 26110, 67668, 100827, 18175, 80987, 95026, 43049, 29875, 73340, + 103134, 127164, 12263, 68500, 44986, 40419, 118430, 83125, 27361, + 83899, 599, 102652, 96223, 98631, 95859, 60979, 40494, 118153, + 71598, 113945, 125768, 2378, 115813, 13443, 7038, 3598, 53377, + 26625, 110573, 3294, 119847, 72218, 85832, 125, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 2023, 29313, 55760, 30415, 117209, 68859, 126088, 105595, 44062, + 129614, 108803, 62812, 119884, 130526, 4554, 108683, 53422, 114202, + 34719, 66446, 65067, 23670, 9591, 82680, 20040, 77609, 34903, + 85401, 37760, 43899, 14395, 67595, 12710, 101806, 77200, 103316, + 65099, 125953, 105217, 67808, 5979, 63677, 12355, 111674, 70131, + 45594, 117232, 10425, 81356, 112691, 14758, 10490, 24556, 27922, + 32980, 20928, 118420, 17204, 4244, 126937, 116849, 106497, 51321, + 114935, 45503, 15461, 59271, 111583, 30113, 103352, 10622, 32510, + 41116, 86928, 10137, 101567, 30707, 124356, 108755, 8203, 11158, + 9603, 114740, 5093, 13054, 61800, 75687, 38080, 11550, 87153, + 33247, 18929, 66437, 13511, 39575, 18765, 61155, 77315, 112366, + 76906, 125693, 40793, 40582, 43161, 81338, 111531, 84813, 49322, + 71309, 83250, 14948, 44745, 13967, 98243, 116072, 5842, 82567, + 77993, 80649, 107659, 66320, 122438, 54394, 54983, 79006, 105681, + 94840, 79085, 41950, 106863, 130420, 89427, 83726, 86511, 44750, + 12837, 47751, 33678, 115313, 66053, 43941, 100068, 107956, 169, + 62673, 70167, 106884, 28460, 37125, 129538, 93387, 56010, 35558, + 40621, 77463, 39765, 16013, 85203, 39465, 19946, 122865, 58068, + 60861, 54036, 48234, 130529, 59321, 83170, 116672, 11733, 94357, + 35207, 46800, 22992, 47306, 80973, 28136, 59828, 4338, 109019, + 16604, 58136, 123247, 75151, 11948, 105570, 86992, 72951, 15873, + 12088, 27387, 60538, 80333, 72819, 97245, 41421, 126078, 82127, + 78747, 70199, 99633, 38926, 56088, 20821, 130509, 125699, 127559, + 47166, 102354, 71946, 79468, 81253, 11120, 37016, 20384, 95174, + 126298, 33756, 60551, 106592, 43758, 59029, 22997, 63753, 65367, + 47855, 5100, 118402, 102326, 108278, 49481, 99388, 114710, 101475, + 124482, 27406, 78309, 52397, 69291, 787, 10, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ========================================================================= + // ou_subhomo 正确性测试 + // + // ou_subhomo(c̃₁, c̃₂) = c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) + // + // 流程: + // 1. 随机生成 BATCH 对密文 (c̃₁, c̃₂),满足 0 < c < n(视作蒙哥马利域值) + // 2. H2D 传输,计时 + // 3. GPU:ou_subhomo → d_result + // 4. D2H 结果,计时 + // 5. 写入 ou_subhomo_test.txt 供 Python 验证 + // ========================================================================= + { + const int BATCH = 200000; + + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + const uint64_t MASK17_T = (1ULL << BASE_BITS) - 1ULL; + const uint64_t CT_TOP = Modn[240]; // = 247,密文最高 limb 严格上界 + + uint64_t *h_c1 = (uint64_t *)malloc(ct_bytes); + uint64_t *h_c2 = (uint64_t *)malloc(ct_bytes); + uint64_t *h_result = (uint64_t *)malloc(ct_bytes); + + srand((unsigned int)time(nullptr)); + + // ── Step 1:随机生成密文对(每个值 < n,视作蒙哥马利域中的 c̃) ────────── + for (int b = 0; b < BATCH; b++) { + uint64_t *c1 = h_c1 + (size_t)b * ARR_LEN; + uint64_t *c2 = h_c2 + (size_t)b * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < 240) { + c1[j] = (((uint64_t)(uint32_t)rand()) ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + c2[j] = (((uint64_t)(uint32_t)rand()) ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17_T; + } else if (j == 240) { + c1[j] = (uint64_t)rand() % CT_TOP; + c2[j] = (uint64_t)rand() % CT_TOP; + } else { + c1[j] = 0ULL; + c2[j] = 0ULL; + } + } + } + + // ── Step 2:设备端分配,H2D 计时 ───────────────────────────────────────── + uint64_t *d_c1 = nullptr, *d_c2 = nullptr, *d_result_sub = nullptr; + CUDA_CHECK(cudaMalloc(&d_c1, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_result_sub, ct_bytes)); + + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_c2, h_c2, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_h2d = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_h2d, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 3:GPU ou_subhomo,计时 ───────────────────────────────────────── + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + ou_subhomo(BATCH, d_c1, d_c2, d_result_sub, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, + d_ctR, // R² mod n,CT 形式 + d_nctR, // R² mod n,NCT 形式 + Modn, // OU 模数 n(主机端,非 n²) + false); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_gpu = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_gpu, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 4:D2H,计时 ──────────────────────────────────────────────────── + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_result, d_result_sub, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float ms_d2h = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms_d2h, ev0, ev1)); + CUDA_CHECK(cudaEventDestroy(ev0)); + CUDA_CHECK(cudaEventDestroy(ev1)); + + // ── Step 5:打印计时 + // ────────────────────────────────────────────────────── + printf("\n========== ou_subhomo test ==========\n"); + printf(" batch=%d ARR_LEN=%d BASE_BITS=%d\n", BATCH, ARR_LEN, BASE_BITS); + printf(" [H2D c1+c2] : %8.4f ms (%6.4f ms/op)\n", ms_h2d, + ms_h2d / BATCH); + printf(" [GPU ou_subhomo] : %8.4f ms (%6.4f ms/op)\n", ms_gpu, + ms_gpu / BATCH); + printf(" [D2H result] : %8.4f ms (%6.4f ms/op)\n", ms_d2h, + ms_d2h / BATCH); + printf("=====================================\n\n"); + /* + // ── Step 6:写测试数据文件 + ───────────────────────────────────────────────── + // 格式 ou_subhomo_test.txt: + // 行 1 (注释):元信息 + // 行 2 (N:) :OU 模数 n(256 limb,base-2^17) + // 每用例 3 行 C1: / C2: / R: + FILE *fp = fopen("ou_subhomo_test.txt", "w"); + if (!fp) { + fprintf(stderr, "[ERROR] 无法创建 ou_subhomo_test.txt\n"); + } else { + fprintf(fp, + "# ou_subhomo test BATCH=%d ARR_LEN=%d BASE_BITS=%d + CT_BASE=2^17\n", BATCH, ARR_LEN, BASE_BITS); fprintf(fp, "N:"); for (int j = + 0; j < ARR_LEN; j++) fprintf(fp, " %llu", (unsigned long long)Modn[j]); + fprintf(fp, "\n"); + for (int b = 0; b < BATCH; b++) { + const uint64_t *c1 = h_c1 + (size_t)b * ARR_LEN; + const uint64_t *c2 = h_c2 + (size_t)b * ARR_LEN; + const uint64_t *r = h_result + (size_t)b * ARR_LEN; + fprintf(fp, "C1:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c1[j]); + fprintf(fp, "\nC2:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)c2[j]); + fprintf(fp, "\nR:"); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fp, " %llu", (unsigned long long)r[j]); + fprintf(fp, "\n"); + } + fclose(fp); + printf("[ou_subhomo test] 测试数据已写入 ou_subhomo_test.txt\n"); + } + */ + // ── 释放 + // ────────────────────────────────────────────────────────────────── + free(h_c1); + free(h_c2); + free(h_result); + cudaFree(d_c1); + cudaFree(d_c2); + cudaFree(d_result_sub); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_addhomo.cu b/heu/library/algorithms/paillier_new/paillier_addhomo.cu new file mode 100644 index 0000000..9bb881c --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_addhomo.cu @@ -0,0 +1,18804 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ============================================================================= +// paillier_addhomo: 密文同态加法 c̃ = FMLM(c̃₁, c̃₂) = c₁c₂R (mod n²) +// +// 数学背景(图示公式): +// AddHomo(c₁, c₂) = c₁·c₂ mod n² +// 在蒙哥马利域中等价于: +// c̃₁ = c₁·R, c̃₂ = c₂·R +// c̃ = FMLM(c̃₁, c̃₂) = c̃₁·c̃₂·R⁻¹ = c₁c₂R (即 (c₁c₂) 的蒙哥马利表示) +// +// 实现:一次 XYfixWarpVector 内核调用,全程在 GPU 上完成,无 H2D/D2H 传输。 +// +// 参数: +// batch 多项式对数(= 待处理的密文对数) +// d_c1_tilde [batch×ARR_LEN] 输入 c̃₁,设备指针(只读) +// d_c2_tilde [batch×ARR_LEN] 输入 c̃₂,设备指针(只读,不被修改) +// d_result [batch×ARR_LEN] 输出 c̃ = FMLM(c̃₁, c̃₂),设备指针 +// 允许 d_result == d_c1_tilde(原地),但不能等于 d_c2_tilde +// 其余参数 与 paillier_subhomo 保持一致(NTT 旋转因子、模数等) +// ============================================================================= +void paillier_addhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // 把 c̃₁ 加载进输出缓冲区;XYfixWarpVector 将原地把它替换为 FMLM(c̃₁, c̃₂) + if (d_result != d_c1_tilde) { + CUDA_CHECK(cudaMemcpy(d_result, d_c1_tilde, batch_bytes, + cudaMemcpyDeviceToDevice)); + } + + // 每个 block = 1 个 warp (32 线程),处理 1 个多项式;共 batch 个 block + // d_c2_tilde 在内核中仅被读取,const_cast 是安全的 + XYfixWarpVector<<>>( + d_result, const_cast(d_c2_tilde), d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + /* + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 1000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + n²))──────────────────────────────── srand((unsigned int)time(NULL)); for + (int p = 0; p < NUM_DEC; p++) { uint64_t *c = h_ct + (size_t)p * ARR_LEN; for + (int j = 0; j < ARR_LEN; j++) { if (j < n2_top_idx) { c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & mask17; } + else if (j == n2_top_idx) { c[j] = (uint64_t)(rand() % (int)n2_top_val); } + else { c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = ((uint64_t)(uint32_t)rand() << 32) | + (uint64_t)(uint32_t)rand(); for (int p = 0; p < NUM_DEC; p++) memcpy(h_exp + + (size_t)p * EXP_U64_LIMBS, lambda_limbs, EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec ) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, + ARR_LEN, BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", + r, round_ms[r], round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + */ + // ============================================================ + // [paillier_addhomo 测试] + // GPU 密文同态加法:c̃ = FMLM(c̃₁, c̃₂) = c₁c₂R (mod n²) + // + // 流程: + // ① CPU 生成随机 c₁, c₂ ∈ (0, n²)(标准域,约 4095 bit) + // ② GPU XYfixWarpROneVector:c₁,c₂ → c̃₁,c̃₂ (标准域 → FMLM 域) + // ③ paillier_addhomo(c̃₁, c̃₂) → c̃_result (纯 GPU kernel) + // ④ 结果写入 addhomo_results.txt,由 verify_addhomo.py 验证 + // ============================================================ + { + const int NUM_ADD = 200000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t add_bytes = (size_t)NUM_ADD * ARR_LEN * sizeof(uint64_t); + + printf("\n============================================================\n"); + printf(" [paillier_addhomo] GPU 密文同态加法测试\n"); + printf(" 批大小=%d,密文约 4095 bit(< n² ≈ 2^4095)\n", NUM_ADD); + printf("============================================================\n"); + + // ── 1. CPU 随机生成 c₁, c₂ ∈ (0, n²)(标准域,约 4095 bit)──────────── + uint64_t *h_c1_std = (uint64_t *)malloc(add_bytes); + uint64_t *h_c2_std = (uint64_t *)malloc(add_bytes); + uint64_t *h_result = (uint64_t *)malloc(add_bytes); + + if (!h_c1_std || !h_c2_std || !h_result) { + fprintf(stderr, "[AddHomo] host malloc 失败,跳过测试\n"); + free(h_c1_std); + free(h_c2_std); + free(h_result); + goto add_done; + } + + srand((unsigned int)time(NULL)); + { + const uint64_t MASK17 = (1ULL << BASE_BITS) - 1; + for (int p = 0; p < NUM_ADD; p++) { + uint64_t *c1p = h_c1_std + (size_t)p * ARR_LEN; + uint64_t *c2p = h_c2_std + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < N2_TOP_IDX) { + c1p[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17; + c2p[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17; + } else if (j == N2_TOP_IDX) { + c1p[j] = 1 + (uint64_t)(rand() % (int)(N2_TOP_VAL - 1)); + c2p[j] = 1 + (uint64_t)(rand() % (int)(N2_TOP_VAL - 1)); + } else { + c1p[j] = 0ULL; + c2p[j] = 0ULL; + } + } + } + } + printf("[AddHomo] 随机 c₁, c₂ 生成完成(%d 对,各约 4095 bit)\n", NUM_ADD); + + // ── 2. 分配 GPU 缓冲区 ─────────────────────────────────────────────────── + { + uint64_t *d_c1_std = NULL; + uint64_t *d_c2_std = NULL; + uint64_t *d_c1_tilde = NULL; + uint64_t *d_c2_tilde = NULL; + uint64_t *d_add_res = NULL; + uint64_t *d_ctR_bat = NULL; + uint64_t *d_nctR_bat = NULL; + + CUDA_CHECK(cudaMalloc(&d_c1_std, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_std, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_c1_tilde, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_tilde, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_add_res, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_ctR_bat, add_bytes)); + CUDA_CHECK(cudaMalloc(&d_nctR_bat, add_bytes)); + + // ── H2D(计时)──────────────────────────────────────────────────── + float ms_h2d = 0.f; + { + cudaEvent_t ev0, ev1; + cudaEventCreate(&ev0); + cudaEventCreate(&ev1); + cudaEventRecord(ev0); + CUDA_CHECK( + cudaMemcpy(d_c1_std, h_c1_std, add_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_c2_std, h_c2_std, add_bytes, cudaMemcpyHostToDevice)); + cudaEventRecord(ev1); + cudaEventSynchronize(ev1); + cudaEventElapsedTime(&ms_h2d, ev0, ev1); + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + printf( + "[AddHomo] H2D(c₁+c₂ → GPU): %.3f ms" + " 数据量=%.2f MB 带宽=%.2f GB/s\n", + ms_h2d, 2.0 * add_bytes / 1048576.0, + 2.0 * add_bytes / 1e9 / (ms_h2d * 1e-3)); + } + + // ── 3. 广播 R²(CT/NCT)到 NUM_ADD 份 ──────────────────────────── + { + const int total = NUM_ADD * ARR_LEN; + const int blk = 256, grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_ctR_bat, d_ctR, NUM_ADD); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_nctR_bat, d_nctR, NUM_ADD); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── 4. 标准域 → FMLM 域:c̃ = c·R mod n² ───────────────────────── + CUDA_CHECK(cudaMemcpy(d_c1_tilde, d_c1_std, add_bytes, + cudaMemcpyDeviceToDevice)); + CUDA_CHECK(cudaMemcpy(d_c2_tilde, d_c2_std, add_bytes, + cudaMemcpyDeviceToDevice)); + + XYfixWarpROneVector<<>>( + d_c1_tilde, d_ctR_bat, d_nctR_bat, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaGetLastError()); + XYfixWarpROneVector<<>>( + d_c2_tilde, d_ctR_bat, d_nctR_bat, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + printf("[AddHomo] c₁, c₂ 已转为 FMLM 域(c̃₁, c̃₂)\n"); + + // ── 5. 预热 ─────────────────────────────────────────────────────── + printf("[AddHomo] 预热 %d 次...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_addhomo(NUM_ADD, d_c1_tilde, d_c2_tilde, d_add_res, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + } + printf("[AddHomo] 预热完成\n\n"); + + // ── 6. cudaEvent 计时(kRounds 轮)────────────────────────────── + printf("[AddHomo] 开始计时(%d 轮 × %d 对密文)...\n", kRounds, NUM_ADD); + float add_ms[10] = {}; + { + cudaEvent_t ev_start, ev_stop; + cudaEventCreate(&ev_start); + cudaEventCreate(&ev_stop); + for (int rnd = 0; rnd < kRounds; rnd++) { + cudaEventRecord(ev_start); + paillier_addhomo( + NUM_ADD, d_c1_tilde, d_c2_tilde, d_add_res, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + cudaEventRecord(ev_stop); + cudaEventSynchronize(ev_stop); + cudaEventElapsedTime(&add_ms[rnd], ev_start, ev_stop); + } + cudaEventDestroy(ev_start); + cudaEventDestroy(ev_stop); + } + printf("[AddHomo] 计时完成\n\n"); + + // ── 7. D2H 最后一轮结果(计时)──────────────────────────────────── + float ms_d2h = 0.f; + { + cudaEvent_t ev0, ev1; + cudaEventCreate(&ev0); + cudaEventCreate(&ev1); + cudaEventRecord(ev0); + CUDA_CHECK( + cudaMemcpy(h_result, d_add_res, add_bytes, cudaMemcpyDeviceToHost)); + cudaEventRecord(ev1); + cudaEventSynchronize(ev1); + cudaEventElapsedTime(&ms_d2h, ev0, ev1); + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + printf( + "[AddHomo] D2H(result → CPU): %.3f ms" + " 数据量=%.2f MB 带宽=%.2f GB/s\n", + ms_d2h, (double)add_bytes / 1048576.0, + (double)add_bytes / 1e9 / (ms_d2h * 1e-3)); + } + /* + // ── 8. 写入 addhomo_results.txt ─────────────────────────────────── + // 文件格式(供 verify_addhomo.py 验证): + // 行1: NUM_ADD ARR_LEN BASE_BITS + // 行2: n² limbs(ARR_LEN 个 uint64,base-2^17,小端序,空格分隔) + // 每个用例 3 行: + // c₁ limbs(标准域) + // c₂ limbs(标准域) + // result limbs(FMLM 域 = c₁c₂R mod n²) + { + FILE *fout = fopen("addhomo_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[AddHomo] 无法创建 addhomo_results.txt\n"); + } else { + fprintf(fout, "%d %d %d\n", NUM_ADD, ARR_LEN, BASE_BITS); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int p = 0; p < NUM_ADD; p++) { + const uint64_t *c1p = h_c1_std + (size_t)p * ARR_LEN; + const uint64_t *c2p = h_c2_std + (size_t)p * ARR_LEN; + const uint64_t *rsp = h_result + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)c1p[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)c2p[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)rsp[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + } + fclose(fout); + printf("[AddHomo] 结果已写入 addhomo_results.txt(%d 条用例)\n", + NUM_ADD); + } + } + */ + // ── 9. 性能统计 ─────────────────────────────────────────────────── + { + float rmin = add_ms[0], rmax = add_ms[0], rsum = 0.f; + for (int r = 0; r < kRounds; r++) { + rsum += add_ms[r]; + if (add_ms[r] < rmin) rmin = add_ms[r]; + if (add_ms[r] > rmax) rmax = add_ms[r]; + } + float ravg = rsum / kRounds; + float e2e_avg = ms_h2d + ravg + ms_d2h; + + printf( + "\n============================================================\n"); + printf(" paillier_addhomo 性能报告\n"); + printf(" 批大小=%d 密文约 4095 bit\n", NUM_ADD); + printf( + "============================================================\n"); + + printf("[数据传输(cudaMemcpy,仅测一次)]\n"); + printf( + "------------------------------------------------------------\n"); + printf( + " H2D(c₁+c₂ → GPU) : %8.3f ms" + " (%.3f us/对 带宽=%.2f GB/s)\n", + ms_h2d, ms_h2d * 1e3f / NUM_ADD, + 2.0 * add_bytes / 1e9 / (ms_h2d * 1e-3)); + printf( + " D2H(result → CPU) : %8.3f ms" + " (%.3f us/对 带宽=%.2f GB/s)\n", + ms_d2h, ms_d2h * 1e3f / NUM_ADD, + (double)add_bytes / 1e9 / (ms_d2h * 1e-3)); + printf(" 传输小计 : %8.3f ms\n", ms_h2d + ms_d2h); + printf( + "------------------------------------------------------------\n"); + + printf("[GPU 计算(cudaEvent,%d 轮,不含传输)]\n", kRounds); + printf( + "------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %8.3f ms (%.3f us/对)\n", r, add_ms[r], + add_ms[r] * 1e3f / NUM_ADD); + printf( + "------------------------------------------------------------\n"); + printf(" GPU 计算平均 : %8.3f ms (%.3f us/对)\n", ravg, + ravg * 1e3f / NUM_ADD); + printf(" GPU 计算最快 : %8.3f ms (%.3f us/对)\n", rmin, + rmin * 1e3f / NUM_ADD); + printf(" GPU 计算最慢 : %8.3f ms (%.3f us/对)\n", rmax, + rmax * 1e3f / NUM_ADD); + printf( + "------------------------------------------------------------\n"); + + printf("[端到端总耗时(H2D + GPU计算 + D2H)]\n"); + printf( + "------------------------------------------------------------\n"); + printf(" 端到端平均 : %8.3f ms (%.3f us/对)\n", e2e_avg, + e2e_avg * 1e3f / NUM_ADD); + printf(" 其中 H2D : %8.3f ms 占比 %.1f%%\n", ms_h2d, + ms_h2d / e2e_avg * 100.f); + printf(" 其中 GPU 计算 : %8.3f ms 占比 %.1f%%\n", ravg, + ravg / e2e_avg * 100.f); + printf(" 其中 D2H : %8.3f ms 占比 %.1f%%\n", ms_d2h, + ms_d2h / e2e_avg * 100.f); + printf( + "============================================================\n"); + } + + // ── 10. 释放 GPU 缓冲区 ─────────────────────────────────────────── + cudaFree(d_c1_std); + cudaFree(d_c2_std); + cudaFree(d_c1_tilde); + cudaFree(d_c2_tilde); + cudaFree(d_add_res); + cudaFree(d_ctR_bat); + cudaFree(d_nctR_bat); + } + + free(h_c1_std); + free(h_c2_std); + free(h_result); + printf("[AddHomo] 测试完成\n"); + add_done:; + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_addhomo2 测试完成。\n"); + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_caddandsubp.cu b/heu/library/algorithms/paillier_new/paillier_caddandsubp.cu new file mode 100644 index 0000000..56a4c3a --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_caddandsubp.cu @@ -0,0 +1,16959 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +/// #include +#include + +#include +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 1024 // 指数 bit 数 +#define EXP_U64_LIMBS (TAU / 64) // = 16,压缩格式每个指数占 16 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 10000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + // ============================================================ + // 1. 分配锁页主机内存(pinned) + // ============================================================ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + uint64_t *h_c1_tilde = NULL; + uint64_t *h_m_batch = NULL; + uint64_t *h_result = NULL; + + CUDA_CHECK(cudaMallocHost(&h_c1_tilde, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_m_batch, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, batch_bytes)); + + // ============================================================ + // 2. 随机填充测试数据 + // c1_tilde:[0, n²) 范围内的随机蒙哥马利密文(limb 表示) + // m_batch :小整数明文,仅低几个 limb 非零 + // ============================================================ + srand((unsigned int)time(NULL)); + + // n² 最高有效 limb 的索引和值(与 save22.cu 中的 n2_top_idx 一致) + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + for (int p = 0; p < NUM_TESTS; p++) { + uint64_t *c = h_c1_tilde + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) + c[j] = ((uint64_t)rand() ^ ((uint64_t)rand() << 15)) & mask17; + else if (j == n2_top_idx) + c[j] = (uint64_t)rand() % n2_top_val; + else + c[j] = 0ULL; + } + } + + for (int p = 0; p < NUM_TESTS; p++) { + uint64_t *m = h_m_batch + (size_t)p * ARR_LEN; + // 明文只填低 4 个 limb,其余为 0(保证 m < n) + for (int j = 0; j < ARR_LEN; j++) m[j] = 0; + m[0] = (uint64_t)rand() & mask17; + m[1] = (uint64_t)rand() & mask17; + m[2] = (uint64_t)rand() & mask17; + m[3] = (uint64_t)rand() % 100; + } + + // ============================================================ + // 3. 分配设备内存(H2D 目标 / D2H 源) + // d_c1:密文输入,Step3 结果也写回此处 + // d_m :明文批 + // ============================================================ + uint64_t *d_c1_dev = NULL; + uint64_t *d_m_dev = NULL; + + CUDA_CHECK(cudaMalloc(&d_c1_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_dev, batch_bytes)); + + // ============================================================ + // 4. [外部] H2D:上传 c̃1 和 m 批,单独计时 + // ============================================================ + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + + printf("[INFO] 开始 H2D 上传(NUM_TESTS=%d, ARR_LEN=%d)...\n", NUM_TESTS, + ARR_LEN); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(d_c1_dev, h_c1_tilde, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_m_dev, h_m_batch, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + + float h2d_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&h2d_ms, ev0, ev1)); + printf("[INFO] H2D 完成:%.4f ms (%.2f GB/s)\n", h2d_ms, + 2.0 * batch_bytes / (h2d_ms * 1e-3) / 1e9); + + // ============================================================ + // 5. 调用修改后的函数(仅 GPU 计算,无 H2D / D2H) + // ============================================================ + printf("[INFO] 执行 GPU 计算(Step1 + Step2 + Step3)...\n"); + StepTime t = paillier_addplain_subplain_batch( + d_c1_dev, d_m_dev, d_N, d_N2, d_ctR, d_nctR, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, + +1 // AddPlain;改为 -1 测试 SubPlain + ); + + // ============================================================ + // 6. [外部] D2H:从 d_c1(Step3 已写入结果)回传到 h_result + // ============================================================ + printf("[INFO] 开始 D2H 下载...\n"); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_result, d_c1_dev, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + + float d2h_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&d2h_ms, ev0, ev1)); + printf("[INFO] D2H 完成:%.4f ms (%.2f GB/s)\n", d2h_ms, + (double)batch_bytes / (d2h_ms * 1e-3) / 1e9); + + // ============================================================ + // 7. 打印完整计时报告 + // ============================================================ + print_steptime_ext("AddPlain batch(H2D/D2H 外部)", t, h2d_ms, d2h_ms); + + // ============================================================ + // 8. 释放资源 + // ============================================================ + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + + cudaFreeHost(h_c1_tilde); + cudaFreeHost(h_m_batch); + cudaFreeHost(h_result); + + cudaFree(d_c1_dev); + cudaFree(d_m_dev); + + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] 测试完成。\n"); + return 0; + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_dec.cu b/heu/library/algorithms/paillier_new/paillier_dec.cu new file mode 100644 index 0000000..2469f73 --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_dec.cu @@ -0,0 +1,18353 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + /* + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 1000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + n²))──────────────────────────────── srand((unsigned int)time(NULL)); for + (int p = 0; p < NUM_DEC; p++) { uint64_t *c = h_ct + (size_t)p * ARR_LEN; for + (int j = 0; j < ARR_LEN; j++) { if (j < n2_top_idx) { c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & mask17; } + else if (j == n2_top_idx) { c[j] = (uint64_t)(rand() % (int)n2_top_val); } + else { c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = ((uint64_t)(uint32_t)rand() << 32) | + (uint64_t)(uint32_t)rand(); for (int p = 0; p < NUM_DEC; p++) memcpy(h_exp + + (size_t)p * EXP_U64_LIMBS, lambda_limbs, EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec ) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, + ARR_LEN, BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", + r, round_ms[r], round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + */ + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 100; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + // n²))──────────────────────────────── + srand((unsigned int)time(NULL)); + for (int p = 0; p < NUM_DEC; p++) { + uint64_t *c = h_ct + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + c[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = + ((uint64_t)(uint32_t)rand() << 32) | (uint64_t)(uint32_t)rand(); + for (int p = 0; p < NUM_DEC; p++) + memcpy(h_exp + (size_t)p * EXP_U64_LIMBS, lambda_limbs, + EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + double h2d_ms = 0.0; + double d2h_ms = 0.0; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + { + struct timespec t0_h2d, t1_h2d; + clock_gettime(CLOCK_MONOTONIC, &t0_h2d); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + clock_gettime(CLOCK_MONOTONIC, &t1_h2d); + h2d_ms = (double)(t1_h2d.tv_sec - t0_h2d.tv_sec) * 1.0e3 + + (double)(t1_h2d.tv_nsec - t0_h2d.tv_nsec) * 1.0e-6; + printf( + "[Dec] H2D 传输: %.3f ms 密文 %.2f MB + 指数 %.2f MB = %.2f MB 带宽 " + "%.2f GB/s\n", + h2d_ms, ct_bytes / 1048576.0, exp_bytes / 1048576.0, + (ct_bytes + exp_bytes) / 1048576.0, + h2d_ms > 0.0 ? (ct_bytes + exp_bytes) / 1.0e6 / h2d_ms : 0.0); + } + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec(d_ct, d_exp_d, d_result, NUM_DEC, TAU, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec(d_ct, d_exp_d, d_result, NUM_DEC, TAU, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + { + struct timespec t0_d2h, t1_d2h; + clock_gettime(CLOCK_MONOTONIC, &t0_d2h); + CUDA_CHECK( + cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + clock_gettime(CLOCK_MONOTONIC, &t1_d2h); + d2h_ms = (double)(t1_d2h.tv_sec - t0_d2h.tv_sec) * 1.0e3 + + (double)(t1_d2h.tv_nsec - t0_d2h.tv_nsec) * 1.0e-6; + printf("[Dec] D2H 传输: %.3f ms 结果 %.2f MB 带宽 %.2f GB/s\n", d2h_ms, + ct_bytes / 1048576.0, + d2h_ms > 0.0 ? ct_bytes / 1.0e6 / d2h_ms : 0.0); + } + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)r_0[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int j = 0; j < EXP_U64_LIMBS; j++) + fprintf(fout, "%llu%c", (unsigned long long)lambda_limbs[j], + j + 1 < EXP_U64_LIMBS ? ' ' : '\n'); + for (int p = 0; p < NUM_DEC; p++) { + const uint64_t *ctilde = h_ct + (size_t)p * ARR_LEN; + const uint64_t *res = h_result + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)ctilde[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)res[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + } + fclose(fout); + printf("[Dec] 结果已写入 dec_results.txt\n"); + } + } + + // ── 9. 统计输出 ─────────────────────────────────────────────────────────── + double sum_ms = 0.0, min_ms = round_ms[0], max_ms = round_ms[0]; + for (int r = 0; r < kRounds; r++) { + sum_ms += round_ms[r]; + if (round_ms[r] < min_ms) min_ms = round_ms[r]; + if (round_ms[r] > max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); + printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, ARR_LEN, + BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时(GPU 核函数,不含传输)]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", r, round_ms[r], + round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("\n[数据传输耗时(单次,pinned 内存)]\n"); + printf("------------------------------------------------------------\n"); + printf(" CPU→GPU (H2D) : %10.3f ms %.2f MB %.2f GB/s\n", h2d_ms, + (ct_bytes + exp_bytes) / 1048576.0, + h2d_ms > 0.0 ? (ct_bytes + exp_bytes) / 1.0e6 / h2d_ms : 0.0); + printf(" GPU→CPU (D2H) : %10.3f ms %.2f MB %.2f GB/s\n", d2h_ms, + ct_bytes / 1048576.0, d2h_ms > 0.0 ? ct_bytes / 1.0e6 / d2h_ms : 0.0); + printf(" 端到端总延迟 : %10.3f ms (H2D + 最快计算轮 + D2H)\n", + h2d_ms + min_ms + d2h_ms); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ============================================================ + // [Getrn 性能测试] + // 计算 hs^{r_i} mod n²(FMLM 域),使用 10-bit 窗口预计算表。 + // 每个测试用例指数 r_i 不同,底数 hs 共享(来自 generate_hs_table)。 + // 结果写入 getrn_results.txt,供 verify_getrn.py 验证正确性。 + // ============================================================ + { + printf("\n============================================================\n"); + printf(" [Getrn] hs^r mod n² 性能测试(10-bit 窗口查表模幂)\n"); + printf("============================================================\n"); + + const int NUM_GETRN = 3000; + const int kWarmupG = 3; + const int kRoundsG = 5; + // Paillier 加密/随机化中 r ← Z_{2^{k/2}},k 为 n 的 bit 长度(约 2048), + // 故 r 最大 k/2 = 1024 bit,比解密用的 λ(2048 bit)减半。 + const int TAU_R = 1024; + const int EXP_R_U64 = TAU_R / 64; // = 16 + const size_t out_bytes_g = (size_t)NUM_GETRN * ARR_LEN * sizeof(uint64_t); + const size_t r_bytes_g = (size_t)NUM_GETRN * EXP_R_U64 * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + + printf("[Getrn] 批大小=%d TAU_R=%d bit WINDOW=%d bit TABLE_SIZE=%d 项\n", + NUM_GETRN, TAU_R, WINDOW_BITS, TABLE_SIZE); + printf("[Getrn] 窗口数=%d(每窗口 %d 次平方 + 1 次乘法)\n", + (TAU_R + WINDOW_BITS - 1) / WINDOW_BITS, WINDOW_BITS); + + // ── 1. CPU 建立预计算表(generate_hs_table)──────────────────────────── + uint64_t h_hs[ARR_LEN] = {}; + uint64_t *h_table = (uint64_t *)malloc(tbl_bytes); + if (!h_table) { + fprintf(stderr, "[Getrn] h_table malloc 失败,跳过\n"); + goto getrn_done; + } + generate_hs_table(N2_arr, h_hs, h_table); + + // ── 2. 随机生成 NUM_GETRN 个 TAU_R-bit 指数 r_i ────────────────────── + { + uint64_t *h_r_g = NULL; + uint64_t *h_res_g = NULL; + CUDA_CHECK(cudaMallocHost(&h_r_g, r_bytes_g)); + CUDA_CHECK(cudaMallocHost(&h_res_g, out_bytes_g)); + + for (int p = 0; p < NUM_GETRN; p++) { + uint64_t *rp = h_r_g + (size_t)p * EXP_R_U64; + for (int j = 0; j < EXP_R_U64; j++) + rp[j] = + ((uint64_t)(uint32_t)rand() << 32) | (uint64_t)(uint32_t)rand(); + } + printf("[Getrn] 随机指数生成完成(%d 个,各 %d bit)\n", NUM_GETRN, + TAU_R); + + // ── 3. 上传表和指数到 GPU(H2D)────────────────────────────────── + uint64_t *d_g_table = NULL; + uint64_t *d_g_r = NULL; + uint64_t *d_g_output = NULL; + CUDA_CHECK(cudaMalloc(&d_g_table, tbl_bytes)); + CUDA_CHECK(cudaMalloc(&d_g_r, r_bytes_g)); + CUDA_CHECK(cudaMalloc(&d_g_output, out_bytes_g)); + { + struct timespec t0h, t1h; + clock_gettime(CLOCK_MONOTONIC, &t0h); + CUDA_CHECK( + cudaMemcpy(d_g_table, h_table, tbl_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_g_r, h_r_g, r_bytes_g, cudaMemcpyHostToDevice)); + clock_gettime(CLOCK_MONOTONIC, &t1h); + double h2d_g = (double)(t1h.tv_sec - t0h.tv_sec) * 1e3 + + (double)(t1h.tv_nsec - t0h.tv_nsec) * 1e-6; + printf( + "[Getrn] H2D: %.3f ms 表 %.2f MB + 指数 %.2f MB" + " 带宽 %.2f GB/s\n", + h2d_g, tbl_bytes / 1048576.0, r_bytes_g / 1048576.0, + h2d_g > 0.0 ? (tbl_bytes + r_bytes_g) / 1e6 / h2d_g : 0.0); + } + + // ── 4. 内核启动参数 ────────────────────────────────────────────── + const int g_wpb = WARP_PER_BLK; + const int g_blks = (NUM_GETRN + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + // ── 5. 预热 ────────────────────────────────────────────────────── + printf("[Getrn] 预热 %d 次...\n", kWarmupG); + for (int w = 0; w < kWarmupG; w++) { + Getrn<<>>( + d_g_table, d_g_r, TAU_R, d_g_output, NUM_GETRN, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaDeviceSynchronize()); + } + CUDA_CHECK(cudaGetLastError()); + printf("[Getrn] 预热完成\n\n"); + + // ── 6. cudaEvent 计时 ──────────────────────────────────────────── + printf("[Getrn] 开始计时(%d 轮 × %d 个)...\n", kRoundsG, NUM_GETRN); + float gms[5] = {}; + cudaEvent_t gev0, gev1; + CUDA_CHECK(cudaEventCreate(&gev0)); + CUDA_CHECK(cudaEventCreate(&gev1)); + for (int rnd = 0; rnd < kRoundsG; rnd++) { + CUDA_CHECK(cudaEventRecord(gev0)); + Getrn<<>>( + d_g_table, d_g_r, TAU_R, d_g_output, NUM_GETRN, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaEventRecord(gev1)); + CUDA_CHECK(cudaEventSynchronize(gev1)); + CUDA_CHECK(cudaEventElapsedTime(&gms[rnd], gev0, gev1)); + } + CUDA_CHECK(cudaEventDestroy(gev0)); + CUDA_CHECK(cudaEventDestroy(gev1)); + + // ── 7. D2H 结果 ────────────────────────────────────────────────── + { + struct timespec t0d, t1d; + clock_gettime(CLOCK_MONOTONIC, &t0d); + CUDA_CHECK(cudaMemcpy(h_res_g, d_g_output, out_bytes_g, + cudaMemcpyDeviceToHost)); + clock_gettime(CLOCK_MONOTONIC, &t1d); + double d2h_g = (double)(t1d.tv_sec - t0d.tv_sec) * 1e3 + + (double)(t1d.tv_nsec - t0d.tv_nsec) * 1e-6; + printf("[Getrn] D2H: %.3f ms 结果 %.2f MB 带宽 %.2f GB/s\n", d2h_g, + out_bytes_g / 1048576.0, + d2h_g > 0.0 ? out_bytes_g / 1e6 / d2h_g : 0.0); + } + + // ── 8. 写结果到 getrn_results.txt ──────────────────────────────── + // 格式(供 verify_getrn.py 验证): + // 行1: NUM_GETRN ARR_LEN BASE_BITS TAU_R + // 行2: n² limbs(ARR_LEN 个 uint64,base-2^17) + // 行3: r_param = table[0] = (2^4352-1) mod n²(FMLM 恒等元) + // 行4: hs limbs(ARR_LEN 个 uint64,base-2^17,标准域) + // 每个用例 2 行: + // r_i limbs(TAU_R/64 个 uint64,64-bit 压缩,LSB limb 在前) + // result limbs(ARR_LEN 个 uint64,base-2^17,FMLM 域) + { + FILE *fg = fopen("getrn_results.txt", "w"); + if (!fg) { + fprintf(stderr, "[Getrn] 无法创建 getrn_results.txt\n"); + } else { + fprintf(fg, "%d %d %d %d\n", NUM_GETRN, ARR_LEN, BASE_BITS, TAU_R); + // n² + for (int j = 0; j < ARR_LEN; j++) + fprintf(fg, "%llu%c", (unsigned long long)N2_arr[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // r_param = (2^4352-1) mod n²(= table[0],FMLM 恒等元) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fg, "%llu%c", (unsigned long long)h_table[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // hs(底数,标准域,base-2^17) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fg, "%llu%c", (unsigned long long)h_hs[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // 每个用例 + for (int p = 0; p < NUM_GETRN; p++) { + const uint64_t *rp = h_r_g + (size_t)p * EXP_R_U64; + const uint64_t *res = h_res_g + (size_t)p * ARR_LEN; + // r_i:TAU_R/64 个 uint64(64-bit packed,LSB 在前) + for (int j = 0; j < EXP_R_U64; j++) + fprintf(fg, "%llu%c", (unsigned long long)rp[j], + j + 1 < EXP_R_U64 ? ' ' : '\n'); + // result:ARR_LEN 个 base-2^17 limb(FMLM 域) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fg, "%llu%c", (unsigned long long)res[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + } + fclose(fg); + printf("[Getrn] 结果已写入 getrn_results.txt\n"); + } + } + + // ── 9. 性能统计 ────────────────────────────────────────────────── + float gmin = gms[0], gmax = gms[0], gsum = 0.f; + for (int rnd = 0; rnd < kRoundsG; rnd++) { + gsum += gms[rnd]; + if (gms[rnd] < gmin) gmin = gms[rnd]; + if (gms[rnd] > gmax) gmax = gms[rnd]; + } + float gavg = gsum / kRoundsG; + + printf( + "\n============================================================\n"); + printf(" Getrn 性能报告\n"); + printf(" 批大小=%d 窗口=%d bit TABLE_SIZE=%d TAU_R=%d bit\n", + NUM_GETRN, WINDOW_BITS, TABLE_SIZE, TAU_R); + printf("============================================================\n"); + printf("[计时结果(cudaEvent)]\n"); + for (int rnd = 0; rnd < kRoundsG; rnd++) + printf(" 轮 %d : %8.3f ms (%.3f us/次)\n", rnd, gms[rnd], + gms[rnd] * 1e3f / NUM_GETRN); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %8.3f ms\n", gavg); + printf(" 最快轮 : %8.3f ms\n", gmin); + printf(" 最慢轮 : %8.3f ms\n", gmax); + printf(" 平均单次 Getrn : %8.3f us\n", gavg * 1e3f / NUM_GETRN); + printf(" (最快轮) 单次 : %8.3f us\n", gmin * 1e3f / NUM_GETRN); + printf("============================================================\n"); + + // ── 10. 释放 Getrn 相关内存 ────────────────────────────────────── + cudaFreeHost(h_r_g); + cudaFreeHost(h_res_g); + cudaFree(d_g_table); + cudaFree(d_g_r); + cudaFree(d_g_output); + } + free(h_table); + printf("[Getrn] 测试完成\n"); + getrn_done:; + } + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_dec_complete.cu b/heu/library/algorithms/paillier_new/paillier_dec_complete.cu new file mode 100644 index 0000000..e667bfc --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_dec_complete.cu @@ -0,0 +1,20113 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12136987446689, 18446744039176586930, 1232056513054754999, + 17214687067477031919, 1317057700433787779, 3682430884702288277, + 15882986415657063096, 16011012896311557528, 8652447280173075721, + 17155031227975517127, 7928940728597073791, 18263282264202751986, + 4200739459691157780, 15364333380163716365, 15064161840435848732, + 5604784277595085801, 13981560403499465143, 16257660673404068315, + 12508060745464592244, 16774628103326688117, 11638111910421033137, + 18360046934253239702, 4007199818774482319, 5193240243026326271, + 3545786642292304913, 17722972152037122125, 8945048594928227839, + 10379898658557360677, 3881119294425181073, 3044810993617697393, + 5084983645078571291, 14695567173062093584, 17473400684945370024, + 9738410036764967581, 17711026959265567629, 18141068788288007674, + 8255491582682120696, 17079453470029463982, 18171400188690694830, + 10663422437703987780, 15272321068286153864, 11033557645731484785, + 8423737106566642542, 10395639523633238707, 14574649230590617068, + 15093480602798076724, 2917566484729317702, 17643995361242251453, + 6516209292763779287, 13063949831114821962, 4201016335371983024, + 1893376350448462048, 3049407863655351179, 18252470266430131139, + 4963025100184055947, 543255064080756700, 7828152310840029591, + 4115727416255336578, 3447257127434672946, 16853056161590062523, + 7905306992414739025, 6794274122063459868, 10626053712895000793, + 9400233939388536864, 13118388822927627762, 661667905440984750, + 6208998755380557787, 7891308247899347390, 4050510349254454901, + 3916416872756535396, 6131708554383996456, 15776732322303289294, + 14311500839918222449, 14233667361511255219, 18272108972806661131, + 27619608770386976, 4961119661516569994, 2629809606622924226, + 1113925432200183160, 2692469845706783749, 5933897446725126277, + 18035799342432954241, 2394021639147018960, 11900046989751532352, + 7882818277897854172, 12851822084613044375, 14311704695587728330, + 16563097150200372140, 17842445160176450298, 9183408955713624427, + 9571073846214621606, 1575183253473706405, 12772339462631047653, + 1653058573447797627, 16705072369051661872, 18014120777360366882, + 3789750418522361569, 187209139274696160, 1662955552233250450, + 11225170967404242941, 10939394779422638460, 3302787312779156041, + 3820334995955691538, 1768339469414493843, 4672199415499252778, + 2185526814921045488, 9933337193684306561, 1504094414550387372, + 5966629681178277712, 12225949730547409727, 17673432338406953131, + 13465479856463801806, 4753147174625419895, 9268984940107969007, + 2281725629233708637, 14391646320283808245, 14084424281150230654, + 13132425219220516011, 5888611885868602733, 1022826975493110915, + 3248082738664524108, 15746649673564735099, 6517449334179110655, + 1590189857791082503, 10768013698187785534, 8701326452582307239, + 9218529050647667269, 16831088990708929380, 15739322550621591612, + 6174629654399123631, 16685659504984935175, 17991978275084551900, + 12672209395566627843, 6921953054696074999, 15716921967334468051, + 6100727058800940905, 15635939569069769356, 1701216518551014748, + 10562190985165201158, 5047274994091701289, 2625410557726277669, + 18432051863318769888, 8937437563733411123, 10441494701693323052, + 3034796417334182972, 15467152725700944885, 17496882520894304741, + 15852876630985719027, 1990330787495014988, 7007776065420025505, + 6850931704519709462, 7430783228549825975, 10314767305758270380, + 17196982377620938041, 11098693348820298230, 7091685251322949927, + 733008192486532245, 18333245046153892245, 214266862155131658, + 10336636183501951817, 16455396338990885719, 13471471725475320033, + 14144179436906969024, 3844364335319284445, 15482549956189583916, + 10712272419739246149, 6420843728054218800, 7686418047460955563, + 11437760842063974582, 17207572282235283264, 4076665451757662674, + 18279964374197619535, 16536596398061992307, 7500161048111875436, + 8484702579543258921, 9190316390210203192, 14333092188354961566, + 7511642897037485262, 8483967287571300293, 3316688893089385183, + 6883315500565789275, 6259525033987456248, 11198250740812718604, + 5374450289247265890, 6350233718241796946, 14390410306336209900, + 8934503715878274427, 12033622953925535977, 17938901204786600648, + 733509976466540003, 7391982170190842137, 12433015778620372353, + 16109197149457099542, 814824089794133102, 4438043559576122045, + 11973610470524742849, 13984518587338919546, 14206361113459729362, + 14731440880014164890, 8278984543462424027, 423779159710002963, + 6936153555096048983, 1700557543899921832, 12860329175219323145, + 10663437606389934101, 9748553182269934872, 16578710159219591753, + 7341085004058451694, 11759526223820241801, 16695552790701510333, + 17846297537355126283, 17247412267345655042, 16582014221048585604, + 8583908311099228303, 2778476983353986135, 18006023971318558808, + 4918782308472758491, 9214552401303372732, 9331267252769858928, + 14579047931415935672, 7045410944764600580, 15779879037114946714, + 3931422602073035642, 4186184302509186885, 2012869482419274130, + 10608961789906967925, 14217693622782194145, 12883786658374524437, + 14609344110091901692, 4165606256439126215, 16522972560076187845, + 1795524687565758616, 3378794700734542774, 12590176114786418660, + 1801375332137468452, 8311818736418349192, 15969620762315655636, + 1437024092595596531, 13267672736221842168, 5710962855326450179, + 14452830359144992103, 7473259782192010852, 7926011145333823737, + 8157869718695693552, 967299641618310079, 17278465872247712370, + 5751150651454550720, 11593731756229954078, 5624315426462315263, + 145966454400421106, 2805927649217484656, 4730620151619002963, + 2658644701417970946, 11249194969060406459, 14626914293964308007, + 5594519421229575032}; +const uint64_t con_modn_shoup[256] = { + 5702931397251532689, 15898375259831901215, 5694435747106281216, + 17976714175742591493, 8022937333770627820, 4023256764252194596, + 9899691033980708753, 18386039868592428787, 7958737770399430416, + 16841092813468060638, 473449321363604613, 8945409407734363057, + 2345650212404534205, 3755667094591430393, 15115129362965580909, + 9569829321427719646, 17072610104501589481, 7234428870262449096, + 10431754657114990467, 834148414668231994, 12165858142001761515, + 18229364958113278685, 11124291784774382140, 15586705826370494763, + 910953971955653251, 5189498502824662190, 10067389405945588110, + 4939562865423853509, 5251301469600028992, 16732471774563845569, + 13443634496105296479, 13798448600935976878, 12357111768224527092, + 5404247608171417908, 6576854357062776800, 10196373708923887660, + 12023237951112779276, 8634579195437162186, 11034118896247081944, + 18381269001254211754, 2649749017176910037, 6893369407195362474, + 10518166280251438160, 2729061371024704804, 11837924634040062248, + 6160747334616383699, 11648346972574376355, 6221509773421703045, + 14190051214678301393, 15511541370874800998, 13860580410285842723, + 2264307686397521109, 13740224204801209575, 1939718287799436287, + 4075312589050378968, 6144394936684063855, 17294774344529229434, + 18076422029736405606, 15039839807266269961, 14881221853008208668, + 8663780317297539480, 17935100243085927763, 13003537821885650695, + 16860664576398245057, 1985636881538439865, 1761914505525338892, + 8620397520520613274, 8170264132653686280, 13963917532215793303, + 15103623488986005970, 5706359745350531277, 15564307873417005419, + 3202689767278678753, 3180045274126185695, 5684054638519149311, + 14581484100665660398, 10842258253812601974, 9012433712480626866, + 645519433587272367, 11336218207121908472, 6955096360553894763, + 10388842993208747782, 14948235341379834871, 16619368396583862911, + 14267268535422444308, 2457513547774775537, 13218815478806995691, + 16044599178974805304, 10945421215527920583, 17205762124994071740, + 18158192091307459316, 15257983888063432725, 17460268377972856095, + 659818834656077040, 16537792610139657939, 11031271672592549033, + 673078441198007644, 3178032089135606972, 10216755798104820307, + 5866947670563017191, 16620847970278045538, 5212853270575474817, + 10708615593892899387, 4788807944276548136, 16554369822607103953, + 9010647092816731400, 9763426428044236363, 5843890743796390412, + 8425921664477514971, 1605181246396393135, 4960985022931098677, + 6238885194053161728, 18321281597130120597, 3206172002937900449, + 64508410495316579, 11661068554019486369, 9179048630229581304, + 3202898955775206065, 2411422857419466274, 6339720133185500467, + 6286767114325201388, 1873790188832210922, 10715199996687697803, + 12283320759177958074, 2296626997500428331, 10651783034682550677, + 16074627283521590646, 10423096039307795482, 5165625588884042964, + 4688269226438782214, 6960383105742896443, 7976605301398407208, + 15727432346721639491, 3863888209647755913, 8113060154858168910, + 1908202978328077049, 2285135013463018403, 10390059577088040765, + 18039323086114210534, 1844936810391851759, 1520130010921724499, + 16928194629126105110, 10447191453663834473, 13572096116754739192, + 2754211345689332869, 2819112478600725532, 14390590876423029186, + 13433643960795464015, 10463629614124117557, 8861488344848785905, + 10683436951681744966, 17966838816743975607, 7019424190886803842, + 1315998778417433792, 8279684889673854498, 271903415331819801, + 16565375601034505389, 16457040058592909323, 6621387438112079884, + 4350366240420774089, 15973472494054450303, 2161943072744535809, + 3156233416460687674, 15858440373624275728, 3158171702744474964, + 14259235795634249751, 14823717284807455265, 16067867969064442539, + 2080520824018040052, 2171682982498993430, 17015351193976049715, + 18375224503011327734, 17360990366822600083, 12820307338773108223, + 7705695981437189977, 16261106053845124768, 9788204544272526811, + 2302557471959371594, 14878965945641539912, 4990373883269078497, + 437384329387965208, 52585105917226204, 8225321005344520346, + 1059473724629972224, 15524449426706454631, 4542759348772867625, + 1491583017230181051, 12541022160247644048, 12611271061643933978, + 13079677897827126003, 11633625736845887921, 555429896845350273, + 15047089863700291453, 7309029822395638953, 7565596654367444329, + 13085542940382773468, 8833547208048204895, 8290661667938146601, + 16535410871561731312, 10525017166045281800, 16860636017320959498, + 8949695391450892511, 8879722660286866657, 17878097955900704129, + 13308489691861540770, 3621809093405865255, 5149701431079656637, + 6036626564656265900, 3024963056937416922, 4437709479395378123, + 2583660210148597841, 4493436546582476704, 18110568265718482172, + 5911772558352596125, 15554147603460275136, 15542442634523032399, + 1458201013794907655, 6079715460834048388, 15968358202182730140, + 9753242274898829149, 13183069962339903146, 8471368658423529662, + 11902918513248342536, 9410543075936915743, 14407464883722455978, + 14428762546967072310, 2733899970636555015, 406377127486695405, + 15676006541852197817, 15417027695641754341, 6826171082178781459, + 2033055020890979910, 12402707508202962383, 436729615700972250, + 5785077018869054601, 11687586423887300225, 3697271174529030011, + 720431861641724346, 904695997883853288, 7021328714187819162, + 11294558343416651450, 3815083684039799982, 14837932449102094072, + 8562937812178914440, 15042485800251768573, 12047267810325699320, + 15513256453531071728, 3405973278442192243, 16199029047884768558, + 9397163461942706468, 1230030495818023290, 17596454057491178894, + 4251831050564258048, 10393169346588382776, 15610999018447648777, + 887271735978662932}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define OU_TAU 1408 // OU 明文 bit 上界(= 22*64,覆盖 1363-bit 私钥 p) +#define OU_EXP_LIMBS (OU_TAU / 64) // = 22,每个 OU 明文占 22 个 uint64_t +#define OU_HR_TAU \ + 128 // H^r' 随机指数 bit 数(= kRandomBits3072,与 CPU 参考实现一致) +#define OU_HR_EXP_LIMBS (OU_HR_TAU / 64) // = 2,每个 r' 占 2 个 uint64_t +#define OU_T_TAU 256 // 解密指数 t 的 bit 数(p-1 的大素因子,~256 bit) +#define OU_T_EXP_LIMBS (OU_T_TAU / 64) // = 4,每个 t 占 4 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================================= +// ou_addhomo:OU 批量密文同态加法(原地版) +// +// 数学依据: +// OU 加密:C = G^m * H^r mod n +// 蒙哥马利(FMLM)域表示:C̃ = C * R mod n +// +// 同态加法利用密文乘法实现: +// Enc(a + b) = Enc(a) * Enc(b) mod n +// 在 FMLM 域中等价于: +// FMLM(C̃₁, C̃₂) = (C₁*R) * (C₂*R) * R⁻¹ mod n +// = C₁ * C₂ * R mod n +// = (C₁*C₂) 的 FMLM 域表示 = Enc(a+b) 的 FMLM 域表示 ✓ +// +// 设计:原地操作 +// XYfixWarpVector 的 inout 参数既是第一操作数的读取位置, +// 也是结果的写回位置,两者必须是同一块内存。 +// 因此本函数直接将 d_c1_tilde 作为 inout 传入,结果原地覆写 d_c1_tilde, +// 无需额外的 d_result 缓冲区,也无需任何 D2D 内存复制。 +// +// 若调用方需要保留原始 c1,应在调用前自行完成备份。 +// +// 输入/输出: +// batch 密文对数量 +// d_c1_tilde [batch × ARR_LEN] 输入 C̃₁(FMLM域);调用后被结果 C̃₁₊₂ 覆写 +// d_c2_tilde [batch × ARR_LEN] 输入 C̃₂(FMLM域,只读,不会被修改) +// +// 返回值:GPU 核函数耗时(ms,cudaEvent 精度) +// ============================================================================= +float ou_addhomo( + int batch, + uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入兼输出:C̃₁ → C̃₁₊₂ + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 C̃₂(只读) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + cudaEvent_t ev_start, ev_stop; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_stop)); + CUDA_CHECK(cudaEventRecord(ev_start)); + + // 每个 block = 1 个 warp(32 线程),处理一对 256-limb 密文。 + // FMLM(C̃₁, C̃₂) 的结果原地写回 d_c1_tilde。 + XYfixWarpVector<<>>( + d_c1_tilde, const_cast(d_c2_tilde), d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(ev_stop)); + CUDA_CHECK(cudaEventSynchronize(ev_stop)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_stop)); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_stop)); + return ms; +} + +// ============================================================= +// ou_mulplain +// +// 功能:OU 批量密文乘明文(c̃^k → c^k·R mod n),使用 CUDA 流并行 +// +// 公式: +// c^p mod n = (G^m · H^r)^p mod n ← OU 密文乘明文 p(明文标量) +// 蒙哥马利域:FMLE(c̃, k) = c^k · R mod n (输入/输出均在蒙哥马利域) +// +// 与 paillier_mulplain 的唯一差异: +// · tau = OU_TAU = 1408(= 22×64,覆盖 1363-bit 明文,高位补零) +// · 指数内存步长 = OU_EXP_LIMBS = 22(不是 EXP_U64_LIMBS = 32) +// · 传入 d_modn 应为 OU 的 n(而非 Paillier 的 n²) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文(设备) +// d_exps [batch * OU_EXP_LIMBS] 明文 k:22 个 uint64_t/个, +// LSB-first,base-2^64,高位补零 +// d_output [batch * ARR_LEN] 输出,蒙哥马利域 c^k·R mod n(设备) +// d_r0 [ARR_LEN] R mod n(Montgomery 域中"1"的表示) +// 其余 NTT 参数与 paillier_mulplain 完全一致 +// ============================================================= +MulPlainTiming ou_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * OU_EXP_LIMBS; // ← 22 个 uint64_t/明文 + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, + OU_TAU, // ← 1408 bit,覆盖 1363-bit 明文 + d_output + ct_off, actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================================= +// ou_build_base_table(内部辅助) +// +// 以给定底数 h_base 建立 WINDOW_BITS-bit 窗口模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = base^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 递推关系(普通大整数乘法,非 FMLM 乘法): +// table[0] = r mod n +// table[i] = table[i-1] × base mod n → = base^i · r mod n ✓ +// ============================================================================= +static void ou_build_base_table( + const char *tag, // 日志前缀,如 "[generate_G_table]" + const uint64_t + *Modn, // [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) + const uint64_t *h_base, // [ARR_LEN] 底数(base-2^17 + // 小端序,标准域,只读) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + // ── 1. 转换 n 和 base 为 mp_int + // ────────────────────────────────────────────── + mp_int n_mp, base_mp, cur_mp, r_mp; + mp_init(&n_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(Modn, &n_mp); + bn17_to_mp_u64(h_base, &base_mp); + + // ── 2. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1,再对 n 取模 + // ──────────── r mod n 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 + // 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n_mp, &r_mp); // r_mp = (2^4352 - 1) mod n + + printf("%s FMLM 参数 r mod n bit 长度 = %d\n", tag, mp_count_bits(&r_mp)); + + // ── 3. 建立预计算表:table[i] = base^i · r mod n(FMLM 域)───────────────── + printf("%s 开始建立预计算表(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", tag, + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n(FMLM 域恒等元,即 base^0 · r mod n) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × base mod n(递推得 base^i · r mod n) + // mp_mulmod 支持输出与输入别名(内部创建临时数),原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "%s 预计算表构建完成(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + tag, TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 4. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// generate_G_table +// +// 以 OU 公钥 G 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = G^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_G [ARR_LEN] OU 公钥 G(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_G_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_G, // [ARR_LEN] OU 公钥 G(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_G_table]", Modn, h_G, h_table); +} + +// ============================================================================= +// generate_invG_table +// +// 以 OU 公钥 G 的逆元 invG 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表 +// (FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = invG^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 用途:密文减明文 p 时,计算 c · invG^p mod n = c · G^{-p} mod n。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_invG [ARR_LEN] OU 公钥 invG(base-2^17 +// 小端序,标准域,只读) h_table [TABLE_SIZE × ARR_LEN] +// 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_invG_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t + *h_invG, // [ARR_LEN] OU 公钥 invG(只读,标准域) + uint64_t * + h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +) { + ou_build_base_table("[generate_invG_table]", Modn, h_invG, h_table); +} + +// ============================================================================= +// Getgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 G^{p_i} mod n,输出 FMLM 域结果。 +// +// 算法与 Getrn 完全相同(左到右 WINDOW_BITS-bit 窗口模幂);区别在于: +// · d_G_table 由 generate_G_table 建立:table[i] = G^i · r mod n +// · tau = OU_TAU = 1408(覆盖 1363-bit 私钥 p,对齐到 22×64 bit) +// · 模数 n(OU modulus)通过 NTT 参数隐式传入 +// +// 正确性: +// table[i] = G^i · r mod n(FMLM 域),table[0] = r mod n(FMLM 恒等元) +// 最终 t = G^p · r mod n ∈ FMLM 域 ✓ +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] generate_G_table 输出的 FMLM 域预计算表 +// d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 uint64,每条 22 个 +// limb) tau p 的 bit 长度(传入 OU_TAU = +// 1408) d_output [batch × ARR_LEN] 输出 G^p(FMLM 域,base-2^17) +// batch 明文数量(每个 warp 处理一条) +// +// 共享内存、启动参数与 Getrn 相同: +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getgp<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getgp( + const uint64_t *__restrict__ d_G_table, // [TABLE_SIZE × ARR_LEN] G + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 G^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p(压缩格式:tau/64 个 uint64,小端序) + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](FMLM 域恒等元 r mod n,即 G^0 的 FMLM 表示) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_G_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_G_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 G^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// Getinvgp +// +// GPU 内核:对一批 OU 明文 p_i,计算 invG^{p_i} mod n,输出 FMLM 域结果。 +// +// 与 Getgp 完全对称,区别仅在于使用 invG 预计算表: +// d_invG_table 由 generate_invG_table 建立:table[i] = invG^i · r mod n +// +// 用途:密文减明文 p 时,先用本内核得到 invG^p 的 FMLM 表示, +// 再与密文做一次 FMLM(XYfixWarpVector),即得 c · invG^p mod n。 +// +// 正确性: +// table[i] = invG^i · r mod n,table[0] = r mod n(FMLM 恒等元) +// 最终 t = invG^p · r mod n ∈ FMLM 域 ✓ +// +// 参数:(同 Getgp,d_invG_table 替换 d_G_table) +// d_invG_table [TABLE_SIZE × ARR_LEN] generate_invG_table 输出的 FMLM +// 域预计算表 d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 个 limb) tau p 的 bit +// 长度(传入 OU_TAU = 1408) d_output [batch × ARR_LEN] 输出 +// invG^p(FMLM 域,base-2^17) batch 明文数量 +// ============================================================================= +__global__ void Getinvgp( + const uint64_t *__restrict__ d_invG_table, // [TABLE_SIZE × ARR_LEN] invG + // 预计算表(FMLM 域) + const uint64_t + *__restrict__ d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序) + int tau, // p 的 bit 长度(= OU_TAU = 1408) + uint64_t *d_output, // [batch × ARR_LEN] 输出 invG^p(FMLM 域) + int batch, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // 加载本 warp 的明文指数 p + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_p_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// 初始化 t = table[0](invG^0 的 FMLM 表示 = r mod n) +#pragma unroll + for (int j = 0; j < 8; j++) + t_buf[8 * lane_id + j] = d_invG_table[8 * lane_id + j]; + __syncwarp(); + + // 左到右 WINDOW_BITS-bit 窗口模幂 + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方 +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 WINDOW_BITS-bit 窗口值 + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf + const uint64_t *entry = d_invG_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val]) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// 写回全局内存(t_buf 即 invG^p 的 FMLM 域表示) +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================================= +// AddSubpTiming +// ============================================================================= +struct AddSubpTiming { + float getxp_ms; // Getgp / Getinvgp 内核耗时(ms) + float fmlm_ms; // XYfixWarpVector 内核耗时(ms) + float total_ms; // 端到端耗时(含 cudaMalloc/Free,ms) +}; + +// ============================================================================= +// ou_addAndsubp +// +// 功能:OU 批量密文加/减明文 +// +// 密文加明文(is_add=true): +// c̃_out[i] = FMLM(c̃[i], G^{p_i} · r mod n) +// = c[i] · G^{p_i} · r mod n ← FMLM 域中 c[i]·G^{p_i} mod n +// 的表示 ✓ +// +// 密文减明文(is_add=false): +// c̃_out[i] = FMLM(c̃[i], invG^{p_i} · r mod n) +// = c[i] · invG^{p_i} · r mod n ← FMLM 域中 c[i]·G^{-p_i} mod n +// 的表示 ✓ +// +// 流程: +// Step① Getgp / Getinvgp: +// 窗口模幂查表,计算 G^{p_i}(或 invG^{p_i})的 FMLM 域表示 +// → 写入临时 buffer d_gp_tilde +// Step② XYfixWarpVector: +// FMLM(c̃[i], d_gp_tilde[i]),结果原地覆写 d_ct[i] +// +// 参数: +// is_add true → 密文加明文(使用 d_G_table) +// false → 密文减明文(使用 d_invG_table) +// batch 密文/明文对数量 +// d_ct [batch × ARR_LEN] 输入兼输出密文(FMLM +// 域,原地覆写) d_p_batch [batch × OU_EXP_LIMBS] 明文指数 p(小端序 +// uint64,每条 22 limb) d_G_table [TABLE_SIZE × ARR_LEN] +// generate_G_table 输出的预计算表(FMLM 域) d_invG_table [TABLE_SIZE × +// ARR_LEN] generate_invG_table 输出的预计算表(FMLM 域) 其余为 NTT +// 参数(与 ou_addhomo / Getgp 完全一致) +// +// 返回:AddSubpTiming(Getxp 耗时、FMLM 耗时、端到端耗时,单位 ms) +// ============================================================================= +AddSubpTiming ou_addAndsubp( + bool is_add, int batch, + uint64_t *d_ct, // [batch × ARR_LEN] 输入/输出(FMLM 域,原地) + const uint64_t + *d_p_batch, // [batch × OU_EXP_LIMBS] 明文指数(小端序 uint64) + const uint64_t + *d_G_table, // [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM 域) + const uint64_t + *d_invG_table, // [TABLE_SIZE × ARR_LEN] invG 预计算表(FMLM 域) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + AddSubpTiming t = {0.f, 0.f, 0.f}; + + // ── 分配临时 buffer:存储 G^p(或 invG^p)的 FMLM 域表示 + // ────────────────────── + uint64_t *d_gp_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_gp_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // ── 内核启动参数 + // ────────────────────────────────────────────────────────────── + const int wpb = WARP_PER_BLK; + const int blks = (batch + wpb - 1) / wpb; + const size_t smem = (size_t)wpb * 512 * sizeof(uint64_t); + + // ── cudaEvent 计时 + // ──────────────────────────────────────────────────────────── + cudaEvent_t ev_total_start, ev_total_end; + cudaEvent_t ev_getxp_start, ev_getxp_end; + cudaEvent_t ev_fmlm_start, ev_fmlm_end; + CUDA_CHECK(cudaEventCreate(&ev_total_start)); + CUDA_CHECK(cudaEventCreate(&ev_total_end)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_start)); + CUDA_CHECK(cudaEventCreate(&ev_getxp_end)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_start)); + CUDA_CHECK(cudaEventCreate(&ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_start)); + + // ── Step①:计算 G^p 或 invG^p,写入 d_gp_tilde + // ─────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_getxp_start)); + + if (is_add) { + Getgp<<>>( + d_G_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } else { + Getinvgp<<>>( + d_invG_table, d_p_batch, OU_TAU, d_gp_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + } + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_getxp_end)); + + // ── Step②:FMLM(c̃, G^p̃) → 结果原地写回 d_ct + // ───────────────────────────────── 每个 block = 1 个 warp(32 + // 线程),处理一对密文 XYfixWarpVector(inout=d_ct[i], in=d_gp_tilde[i]) → + // d_ct[i] = c·G^p·r mod n + CUDA_CHECK(cudaEventRecord(ev_fmlm_start)); + XYfixWarpVector<<>>( + d_ct, d_gp_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(ev_fmlm_end)); + + CUDA_CHECK(cudaEventRecord(ev_total_end)); + CUDA_CHECK(cudaEventSynchronize(ev_total_end)); + + // ── 收集耗时 + // ────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventElapsedTime(&t.getxp_ms, ev_getxp_start, ev_getxp_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.fmlm_ms, ev_fmlm_start, ev_fmlm_end)); + CUDA_CHECK(cudaEventElapsedTime(&t.total_ms, ev_total_start, ev_total_end)); + + // ── 清理 + // ────────────────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventDestroy(ev_total_start)); + CUDA_CHECK(cudaEventDestroy(ev_total_end)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_start)); + CUDA_CHECK(cudaEventDestroy(ev_getxp_end)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_start)); + CUDA_CHECK(cudaEventDestroy(ev_fmlm_end)); + cudaFree(d_gp_tilde); + + return t; +} + +// ============================================================================= +// ou_subhomo:OU 密文同态减法(c̃₁ · c̃₂⁻¹ mod n,全 GPU) +// +// 直接复用 paillier_subhomo2 的四步流程,将模数从 Paillier 的 n² +// 替换为 OU 的 n(OU 的 n ≈ 4088 bit < INV_BITS=4096,CGBN 配置兼容): +// +// ① GPU XYfixWarpIRVector : c̃₂ → c₂ (FMLM⁻¹, mod n) +// ②abc GPU CGBN 批量模逆 : c₂ → c₂⁻¹ (mod n) +// ③ GPU XYfixWarpROneVector: c₂⁻¹ → c̃₂⁻¹ (FMLM × R²modn, mod n) +// ④ GPU XYfixWarpVector : c̃₁·c̃₂⁻¹ → c̃ (FMLM, mod n) +// +// 参数: +// d_c1_tilde [batch×ARR_LEN] 被减数密文(蒙哥马利域,只读) +// d_c2_tilde [batch×ARR_LEN] 减数密文 (蒙哥马利域,只读) +// d_result [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ mod n(蒙哥马利域) +// d_negmodn / d_negmodn_shoup NTT 域 -n(CT 形式)及其 Shoup 值 +// d_modn / d_modn_shoup NTT 域 n(NCT 形式)及其 Shoup 值 +// mod NTT 素模数(25-bit,即 MOD) +// d_twiddle…/d_ICTTwiddle…/ +// d_NCTtwiddle…/d_InvNCTtwiddle… 四组旋转因子及 Shoup 值 +// inv_val / inv_shoup_val INTT 缩放因子 N^{-1} mod MOD 及其 Shoup 值 +// d_sample [ARR_LEN] FMLM 采样辅助数组 +// d_r2modn_ct [ARR_LEN] R² mod n 的 CT-NTT 形式 +// (对应 main 中的 d_ctR,由 Testctsample +// 生成) +// d_r2modn_nct [ARR_LEN] R² mod n 的 NCT-NTT 形式 +// (对应 main 中的 d_nctR,由 Testnctsample +// 生成) +// Nn_arr [ARR_LEN] OU 模数 n(主机端,base-2^17 小端序) +// verbose 是否打印各步进度 +// ============================================================================= +void ou_subhomo( + int batch, const uint64_t *d_c1_tilde, const uint64_t *d_c2_tilde, + uint64_t *d_result, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r2modn_ct, // R² mod n,CT-NTT 形式 (main 中的 d_ctR) + const uint64_t *d_r2modn_nct, // R² mod n,NCT-NTT 形式 (main 中的 d_nctR) + const uint64_t *Nn_arr, // OU 模数 n,主机端 base-2^17(非 n²) + bool verbose) { + paillier_subhomo2( + batch, d_c1_tilde, d_c2_tilde, d_result, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample, + d_r2modn_ct, // paillier_subhomo2 此处期望 R² mod n²;OU 传 R² mod n + d_r2modn_nct, + Nn_arr, // paillier_subhomo2 此处期望 n²;OU 传 n + verbose); +} + +// ============================================================================= +// generate_H_table +// +// 以 OU 公钥 H 为底数,建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU +// 端): +// h_table[i × ARR_LEN] = H^i · r mod n, i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN×BASE_BITS) - 1 mod n(FMLM 域恒等元)。 +// H 以 base-2^17 小端序表示(256 limb),与 G/invG 格式完全一致。 +// +// 参数: +// Modn [ARR_LEN] OU 模数 n(base-2^17 小端序,只读) +// h_H [ARR_LEN] OU 公钥 H(base-2^17 小端序,标准域,只读) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_H_table( + const uint64_t *Modn, // [ARR_LEN] OU 模数 n(只读) + const uint64_t *h_H, // [ARR_LEN] OU 公钥 H(只读,标准域) + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + ou_build_base_table("[generate_H_table]", Modn, h_H, h_table); +} + +// ============================================================================= +// ou_gen_r_prime +// +// CPU 端批量生成随机指数 r',与 CPU 参考实现(kRandomBits3072=128 +// bit)保持一致。 +// +// 与参考实现的对应关系: +// encryptor.cc: BigInt r = BigInt::RandomExactBits(random_bits_) +// 其中 n.BitCount() >= 2560 时 random_bits_ = kRandomBits3072 = +// 128 +// +// r' 仅 128 bit,远小于 p ≈ 1363 bit,无需与 p 比较,直接随机填充即可。 +// 安全性来自计算复杂度(128-bit 计算安全),而非统计均匀性。 +// +// 参数: +// batch 生成数量 +// h_r_prime [batch × OU_HR_EXP_LIMBS] 输出随机指数(base-2^64 小端序, +// OU_HR_EXP_LIMBS=2 个 uint64_t) +// ============================================================================= +// 保留参数兼容调用处,128-bit 时不需要 p +static void ou_gen_r_prime(const uint64_t * /*ou_p_limbs17*/, int batch, + uint64_t *h_r_prime) { + // 生成 batch 个 OU_HR_TAU=128 bit 随机数,每个用 OU_HR_EXP_LIMBS=2 个 + // uint64_t 存储。 rand() 在 Windows 下返回 15-bit 值(RAND_MAX=32767), 用 4 + // 次 rand() 拼装 60-bit,再补高位,得到足够随机的 64-bit 值。 + for (int k = 0; k < batch; k++) { + uint64_t *r = h_r_prime + (size_t)k * OU_HR_EXP_LIMBS; + for (int j = 0; j < OU_HR_EXP_LIMBS; j++) + r[j] = ((uint64_t)rand() << 45) | ((uint64_t)rand() << 30) | + ((uint64_t)rand() << 15) | (uint64_t)rand(); + } +} + +// ============================================================================= +// ou_randomize +// +// 功能:批量随机化 OU 密文 +// c̃_new[i] = FMLM(c̃[i], H^{r'_i}) mod n (FMLM 域,原地覆写) +// +// 流程(与 paillier_randomize2 完全对称,模数为 OU 的 n): +// Step① Getrn:用 H 预计算表(d_H_table)计算 H^{r'_i}(FMLM 域) +// Step② XYfixWarpVector:c̃[i] · H^{r'_i} → 原地写回 d_c +// +// 调用前准备: +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_c [batch × ARR_LEN] 输入兼输出密文(FMLM 域) +// d_H_table [TABLE_SIZE × ARR_LEN] H 预计算表(FMLM +// 域,generate_H_table 输出) d_r_prime_batch [batch × OU_HR_EXP_LIMBS] +// 随机指数 r'(OU_HR_TAU=128 bit,base-2^64 小端序,ou_gen_r_prime 生成) +// batch 密文数量 +// +// 返回:Step①+② GPU 总耗时(ms) +// ============================================================================= +float ou_randomize( + uint64_t *d_c, const uint64_t *d_H_table, const uint64_t *d_r_prime_batch, + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 Getrn 输出 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:XYfixWarpVector — c̃[i] · H^{r'_i} mod n(FMLM 域,原地) + XYfixWarpVector<<>>( + d_c, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================================= +// ou_encrypt +// +// 功能:批量 OU 加密(正明文) +// c̃[i] = FMLM(G^{m_i}, H^{r'_i}) mod n (FMLM 域输出) +// +// 数学依据: +// OU 加密公式:c = G^m · H^r mod n +// FMLM 域语义:FMLM(Ã, B̃) = ÷B̃·R⁻¹ mod n +// 因为 Getgp/Getrn 输出 FMLM 域(乘了 R),所以: +// FMLM(G^m·R, H^r'·R) = G^m·H^r'·R mod n ✓ +// +// 三步流程: +// Step① Getgp :G 预计算表 × m_i → G^{m_i} (FMLM 域)→ d_ct_out +// Step② Getrn :H 预计算表 × r'_i → H^{r'_i}(FMLM 域)→ d_Hr(临时) +// Step③ XYfixWarpVector:d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) +// +// 调用前准备: +// generate_G_table(Modn, ou_G, h_G_table) 并上传 → d_G_table +// generate_H_table(Modn, ou_H, h_H_table) 并上传 → d_H_table +// ou_gen_r_prime(ou_p, batch, h_r_prime) 并上传 → d_r_prime_batch +// +// 参数: +// d_G_table [TABLE_SIZE × ARR_LEN] G 预计算表(FMLM +// 域,generate_G_table 输出) d_H_table [TABLE_SIZE × ARR_LEN] H +// 预计算表(FMLM 域,generate_H_table 输出) d_m_batch [batch × +// OU_EXP_LIMBS] 明文 m(base-2^64 小端序,OU_EXP_LIMBS=22) +// d_r_prime_batch [batch × OU_HR_EXP_LIMBS] 随机指数 r'(OU_HR_TAU=128 +// bit,OU_HR_EXP_LIMBS=2) d_ct_out [batch × ARR_LEN] 输出密文(FMLM +// 域) batch 明文数量 +// +// 返回:GPU 端 Step①②③ 总耗时(ms) +// ============================================================================= +float ou_encrypt(const uint64_t *d_G_table, const uint64_t *d_H_table, + const uint64_t *d_m_batch, // [batch × OU_EXP_LIMBS] + const uint64_t *d_r_prime_batch, // [batch × OU_HR_EXP_LIMBS] + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM 域) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 临时 buffer:存储 H^{r'_i}(FMLM 域) + uint64_t *d_Hr = nullptr; + CUDA_CHECK(cudaMalloc(&d_Hr, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getgp — G^{m_i} mod n(FMLM 域)→ d_ct_out + Getgp<<>>( + d_G_table, d_m_batch, OU_TAU, d_ct_out, batch, d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:Getrn — H^{r'_i} mod n(FMLM 域)→ d_Hr + Getrn<<>>( + d_H_table, d_r_prime_batch, OU_HR_TAU, d_Hr, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step③:XYfixWarpVector — d_ct_out · d_Hr → d_ct_out(原地,FMLM 乘法) + // 结果:c̃[i] = G^{m_i} · H^{r'_i} · R mod n(FMLM 域密文) + XYfixWarpVector<<>>( + d_ct_out, d_Hr, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + CUDA_CHECK(cudaFree(d_Hr)); + return ms; +} + +// ============================================================= +// ou_dec +// +// 功能:GPU 批量 OU 解密第一步——模幂 c^t mod n(普通域输出) +// +// OU 完整解密流程: +// Step①(本函数): c^t mod n ← FMLE_mod3_Kernel,NTT 参数为 n 模数 +// Step②(CPU 端): (result) mod p² ← 因 n=p²·q,c^t mod n 再 mod p² = c^t +// mod p² Step③(CPU 端): m = L(c') · gp_inv mod p,其中 L(x) = (x-1)/p +// +// 输入 d_c_tilde 为 FMLM 域密文 c̃ = c·R mod n; +// FMLE_mod3_Kernel 跳过 Step8(输入已是 FMLM 域), +// 保留 Step13(乘 R⁻¹,还原为普通域),输出 c^t mod n。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(FMLM 域,mod n) +// d_t_exp [batch × OU_T_EXP_LIMBS] 指数 t(压缩格式:OU_T_EXP_LIMBS=4 个 +// uint64_t, +// 小端序,位 i 在 limb[i/64] 的第 i%64 +// 位) +// d_output [batch × ARR_LEN] 输出:c^t mod n(普通域,base-2^17) +// batch 批大小 +// d_r0 [ARR_LEN] r₀ = (2^(ARR_LEN×BASE_BITS)-1) mod n +// (FMLM 单位元,即蒙哥马利域中的 1) +// 其余为 NTT 参数(n 模数,与 ou_encrypt 完全一致) +// +// 返回:GPU 端到端耗时(ms) +// ============================================================= +float ou_dec(const uint64_t *d_c_tilde, const uint64_t *d_t_exp, + uint64_t *d_output, int batch, const uint64_t *d_r0, + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * OU_T_EXP_LIMBS; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_t_exp + exp_off, OU_T_TAU, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================================ +// ou_broadcast_kernel: 将单份 src[ARR_LEN] 广播到 dst[batch × ARR_LEN] +// ============================================================================ +__global__ void ou_broadcast_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================================ +// ou_compute_p2_kernel: 单 instance CGBN 计算 p² = p * p +// 启动参数: <<<1, INV_TPI>>> +// ============================================================================ +__global__ void ou_compute_p2_kernel(cgbn_error_report_t *report, + inv_bn_mem_t *d_p2, inv_bn_mem_t *d_p) { + int instance_id = (blockIdx.x * blockDim.x + threadIdx.x) / INV_TPI; + if (instance_id != 0) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t p, p2; + cgbn_load(env, p, d_p); + cgbn_mul(env, p2, p, p); // p ≈ 1364 bit, p² ≈ 2728 bit < 4096 bit,不截断 + cgbn_store(env, d_p2, p2); +} + +// ============================================================================ +// ou_L_kernel: GPU CGBN 批量计算 L = (c^t mod p² − 1) / p +// 每 INV_TPI(=32) 线程处理一个实例 +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_L_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_L, // [batch] 输出 + inv_bn_mem_t *d_ct, // [batch] 输入:c^t(CGBN 格式) + inv_bn_mem_t *d_p2, // [1] 输入:p²(常量,所有实例共用) + inv_bn_mem_t *d_p, // [1] 输入:p (常量,所有实例共用) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t ct, p2, p_bn, tmp, L; + + cgbn_load(env, ct, d_ct + instance_id); + cgbn_load(env, p2, d_p2); // 所有实例共用同一个 p² + cgbn_load(env, p_bn, d_p); + + cgbn_rem(env, tmp, ct, p2); // tmp = c^t mod p² + cgbn_sub_ui32(env, tmp, tmp, + 1); // tmp = tmp − 1 (OU 保证 c^t ≡ 1 mod p,故 tmp ≥ 1) + cgbn_div(env, L, tmp, p_bn); // L = tmp / p (精确整除) + + cgbn_store(env, d_L + instance_id, L); +} + +// ============================================================================ +// ou_modp_kernel: GPU CGBN 批量计算 m = prod mod p +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void ou_modp_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_m, // [batch] 输出 + inv_bn_mem_t *d_prod, // [batch] 输入:L × gp_inv mod n + inv_bn_mem_t *d_p, // [1] 输入:p(常量) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t prod, p_bn, m; + + cgbn_load(env, prod, d_prod + instance_id); + cgbn_load(env, p_bn, d_p); + cgbn_rem(env, m, prod, p_bn); + + cgbn_store(env, d_m + instance_id, m); +} + +// ============================================================================ +// ou_dec_complete: 完整 OU 解密 +// +// Step① ou_dec : d_c_tilde → d_ct_plain (c^t mod n,普通域) +// Step② format_to_cgbn_kernel : d_ct_plain → d_ct_cgbn +// Step③ ou_L_kernel : d_ct_cgbn → d_L_cgbn (L=(c^t mod +// p²−1)/p) Step④ format_from_cgbn_kernel : d_L_cgbn → d_L_b17 Step⑤ +// XYfixWarpROneVector×1 : d_gp_inv → 蒙哥马利域(原地,单份) Step⑥ +// ou_broadcast_kernel : d_gp_inv → d_gp_inv_batch(batch 份) Step⑦ +// XYfixWarpVector : d_L_b17 × d_gp_inv_batch → L×gp_inv mod n Step⑧ +// format_to_cgbn_kernel : d_L_b17 → d_prod_cgbn Step⑨ ou_modp_kernel : +// d_prod_cgbn → d_m_cgbn (mod p) Step⑩ format_from_cgbn_kernel : d_m_cgbn +// → d_m_out +// +// 返回: GPU 全流程耗时(ms,含 ou_dec 内部时间) +// ============================================================================ +float ou_dec_complete( + const uint64_t *d_c_tilde, // [batch × ARR_LEN] FMLM 域密文 + const uint64_t *d_t_exp, // [batch × OU_T_EXP_LIMBS] 指数 t + uint64_t *d_m_out, // [batch × ARR_LEN] 输出:明文 m(base-2^17) + int batch, + const uint64_t *h_p, // p(主机端,base-2^17) + const uint64_t *h_gp_inv, // gp_inv(主机端,base-2^17) + const uint64_t *d_r0, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + uint64_t *d_ctR, // CT(R²),传给 XYfixWarpROneVector + uint64_t *d_nctR // NCT(R²),传给 XYfixWarpROneVector +) { + const size_t ct_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + const size_t cgbn_bytes = (size_t)batch * sizeof(inv_bn_mem_t); + const size_t one_cgbn = sizeof(inv_bn_mem_t); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + // ── 申请设备临时内存 ───────────────────────────────────────────────────── + uint64_t *d_ct_plain = nullptr; + inv_bn_mem_t *d_ct_cgbn = nullptr; + inv_bn_mem_t *d_p_cgbn = nullptr; + inv_bn_mem_t *d_p2_cgbn = nullptr; + inv_bn_mem_t *d_L_cgbn = nullptr; + uint64_t *d_L_b17 = nullptr; + uint64_t *d_gp_inv = nullptr; + uint64_t *d_gp_inv_batch = nullptr; + inv_bn_mem_t *d_prod_cgbn = nullptr; + inv_bn_mem_t *d_m_cgbn = nullptr; + + CUDA_CHECK(cudaMalloc(&d_ct_plain, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_p_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_p2_cgbn, one_cgbn)); + CUDA_CHECK(cudaMalloc(&d_L_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_b17, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_gp_inv, ARR_LEN * sizeof(uint64_t))); + CUDA_CHECK(cudaMalloc(&d_gp_inv_batch, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_prod_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_cgbn, cgbn_bytes)); + + // ── Step①: c^t mod n ──────────────────────────────────────────────────── + ou_dec(d_c_tilde, d_t_exp, d_ct_plain, batch, d_r0, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + // ou_dec 内部使用多流,需同步后才能进行后续 CGBN 操作 + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 上传 p、gp_inv;GPU 计算 p² ────────────────────────────────────────── + { + inv_bn_mem_t h_p_cgbn; + bn17_to_cgbn_mem(h_p, &h_p_cgbn); + CUDA_CHECK( + cudaMemcpy(d_p_cgbn, &h_p_cgbn, one_cgbn, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_gp_inv, h_gp_inv, ARR_LEN * sizeof(uint64_t), + cudaMemcpyHostToDevice)); + + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + ou_compute_p2_kernel<<<1, INV_TPI>>>(report, d_p2_cgbn, d_p_cgbn); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] p² 计算 CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step②: base-2^17 → CGBN ───────────────────────────────────────────── + format_to_cgbn_kernel<<>>(d_ct_cgbn, d_ct_plain, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③: L = (c^t mod p² − 1) / p ──────────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_L_kernel<<>>(report, d_L_cgbn, d_ct_cgbn, d_p2_cgbn, + d_p_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_L_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step④: CGBN → base-2^17 (L) ──────────────────────────────────────── + format_from_cgbn_kernel<<>>(d_L_b17, d_L_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑤: gp_inv 转蒙哥马利域(单份)────────────────────────────────── + // FMLM(gp_inv, R²) = gp_inv * R mod n → 蒙哥马利域,结果原地写回 d_gp_inv + XYfixWarpROneVector<<<1, 32>>>( + d_gp_inv, d_ctR, d_nctR, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑥: 广播 gp_inv_mont 到 batch 份 ──────────────────────────────── + { + const int total = batch * ARR_LEN; + const int bthreads = 256; + const int bblocks = (total + bthreads - 1) / bthreads; + ou_broadcast_kernel<<>>(d_gp_inv_batch, d_gp_inv, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step⑦: FMLM(L_plain, gp_inv_mont) = L × gp_inv mod n(普通域)────── + // inout=d_L_b17(普通域),inoutA=d_gp_inv_batch(蒙哥马利域) + // 结果 = L * gp_inv_mont * R^{-1} = L * gp_inv mod n,写回 d_L_b17 + XYfixWarpVector<<>>( + d_L_b17, d_gp_inv_batch, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑧: base-2^17 → CGBN (L × gp_inv mod n) ───────────────────────── + format_to_cgbn_kernel<<>>(d_prod_cgbn, d_L_b17, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑨: m = (L × gp_inv mod n) mod p ──────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_modp_kernel<<>>(report, d_m_cgbn, d_prod_cgbn, d_p_cgbn, + batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[ou_dec_complete] ou_modp_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step⑩: CGBN → base-2^17 (m,写入 d_m_out) ────────────────────────── + format_from_cgbn_kernel<<>>(d_m_out, d_m_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + // ── 释放临时设备内存 ───────────────────────────────────────────────────── + cudaFree(d_ct_plain); + cudaFree(d_ct_cgbn); + cudaFree(d_p_cgbn); + cudaFree(d_p2_cgbn); + cudaFree(d_L_cgbn); + cudaFree(d_L_b17); + cudaFree(d_gp_inv); + cudaFree(d_gp_inv_batch); + cudaFree(d_prod_cgbn); + cudaFree(d_m_cgbn); + + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================================ +// paillier_L_kernel: 批量计算 L(c^λ) = (c^λ − 1) / n +// 前提:c^λ ≡ 1 (mod n),故 c^λ − 1 精确被 n 整除 +// 启动参数: <<<(batch*INV_TPI+255)/256, 256>>> +// ============================================================================ +__global__ void paillier_L_kernel( + cgbn_error_report_t *report, + inv_bn_mem_t *d_L, // [batch] 输出:L 值(< n) + inv_bn_mem_t *d_ct, // [batch] 输入:c^λ mod n²(普通域,已在 [0,n²)) + inv_bn_mem_t *d_n, // [1] 输入:n(常量,所有实例共用) + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= batch) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t ct, n_bn, tmp, L; + cgbn_load(env, ct, d_ct + instance_id); + cgbn_load(env, n_bn, d_n); + + cgbn_sub_ui32(env, tmp, ct, 1); // tmp = c^λ − 1 + cgbn_div(env, L, tmp, n_bn); // L = (c^λ − 1) / n(精确整除) + + cgbn_store(env, d_L + instance_id, L); +} + +// 主机辅助函数:base-2^17(256 limbs)→ base-2^64(n64 limbs) +static void bn17_to_bn64(const uint64_t *src17, uint64_t *dst64, int n64) { + memset(dst64, 0, (size_t)n64 * sizeof(uint64_t)); + for (int i = 0; i < 256; i++) { + uint64_t val = src17[i] & 0x1FFFFULL; + int bpos = i * 17; + int widx = bpos / 64; + int boff = bpos % 64; + if (widx >= n64) break; + dst64[widx] |= (val << boff); + if (boff + 17 > 64 && widx + 1 < n64) + dst64[widx + 1] |= (val >> (64 - boff)); + } +} + +// ============================================================================ +// paillier_dec_complete: 完整 Paillier 解密 +// +// Paillier 解密公式:m = L(c^λ mod n²) · μ mod n +// L(x) = (x − 1) / n +// +// Step① paillier_dec : d_c_tilde → d_ct_plain (c^λ mod +// n²,普通域) Step② format_to_cgbn_kernel : d_ct_plain → d_ct_cgbn Step③ +// paillier_L_kernel : d_ct_cgbn → d_L_cgbn (L=(c^λ−1)/n) Step④ +// format_from_cgbn_kernel : d_L_cgbn → d_L_b17 Step⑤ +// XYfixWarpROneVector×1 : d_mu → 蒙哥马利域(单份,原地) Step⑥ +// ou_broadcast_kernel : d_mu → d_mu_batch(batch 份) Step⑦ +// XYfixWarpVector : d_L_b17 × d_mu_batch → L×μ mod n²(普通域) +// Step⑧ format_to_cgbn_kernel : d_L_b17 → d_prod_cgbn +// Step⑨ ou_modp_kernel(复用) : d_prod_cgbn → d_m_cgbn (mod n) +// Step⑩ format_from_cgbn_kernel : d_m_cgbn → d_m_out +// +// 参数说明: +// d_c_tilde [batch × ARR_LEN] 输入:FMLM 域密文 c̃ = c · R mod n² +// d_m_out [batch × ARR_LEN] 输出:明文 m(base-2^17) +// tau λ 的 bit 长度(建议 2048) +// h_lambda [256] λ(host,base-2^17) +// h_mu [256] μ = λ⁻¹ mod n(host,base-2^17) +// h_n [256] Paillier 主模数 n(host,base-2^17) +// d_negmodn/d_modn 等 NTT 参数,对应 n²(加密模数) +// d_ctR/d_nctR R² 的 CT/NCT 形式(XYfixWarpROneVector 用) +// ============================================================================ +float paillier_dec_complete( + const uint64_t *d_c_tilde, uint64_t *d_m_out, int batch, int tau, + const uint64_t *d_lambda_exp, // [batch×tau/64] DEVICE, base-2^64 + uint64_t *d_mu_dev, // [ARR_LEN] DEVICE, base-2^17 + inv_bn_mem_t *d_n_cgbn, // [1] DEVICE, CGBN + const uint64_t *d_r0, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, uint64_t mod, + const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample, + uint64_t *d_ctR, uint64_t *d_nctR) { + const size_t ct_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + const size_t cgbn_bytes = (size_t)batch * sizeof(inv_bn_mem_t); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + // ── 申请设备临时内存 ───────────────────────────────────────────────────── + uint64_t *d_ct_plain = nullptr; + inv_bn_mem_t *d_ct_cgbn = nullptr; + inv_bn_mem_t *d_L_cgbn = nullptr; + uint64_t *d_L_b17 = nullptr; + uint64_t *d_mu_batch = nullptr; // μ_mont broadcast(batch 份) + inv_bn_mem_t *d_prod_cgbn = nullptr; + inv_bn_mem_t *d_m_cgbn = nullptr; + + CUDA_CHECK(cudaMalloc(&d_ct_plain, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_L_b17, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_mu_batch, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_prod_cgbn, cgbn_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_cgbn, cgbn_bytes)); + + // ── Step①: c^λ mod n²(普通域)──────────────────────────────────────── + paillier_dec(d_c_tilde, d_lambda_exp, d_ct_plain, batch, tau, d_r0, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②: base-2^17 → CGBN ───────────────────────────────────────────── + format_to_cgbn_kernel<<>>(d_ct_cgbn, d_ct_plain, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③: L = (c^λ − 1) / n ──────────────────────────────────────────── + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + paillier_L_kernel<<>>(report, d_L_cgbn, d_ct_cgbn, + d_n_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[paillier_dec_complete] paillier_L_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step④: CGBN → base-2^17 (L) ──────────────────────────────────────── + format_from_cgbn_kernel<<>>(d_L_b17, d_L_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑤: μ 转蒙哥马利域(单份,原地) + // XYfixWarpROneVector: FMLM(μ, R²) = μ·R mod n² → 蒙哥马利表示 + XYfixWarpROneVector<<<1, 32>>>( + d_mu_dev, d_ctR, d_nctR, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑥: 广播 μ_mont 到 batch 份 ────────────────────────────────────── + { + const int total = batch * ARR_LEN; + const int bthreads = 256; + const int bblocks = (total + bthreads - 1) / bthreads; + ou_broadcast_kernel<<>>(d_mu_batch, d_mu_dev, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step⑦: FMLM(L_plain, μ_mont) = L·μ mod n²(普通域)──────────────── + // FMLM(A, B) = A·B·R⁻¹ mod n² + // A = L(普通域),B = μ·R(蒙哥马利域) + // 结果 = L·μ·R·R⁻¹ = L·μ mod n² + XYfixWarpVector<<>>( + d_L_b17, d_mu_batch, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑧: base-2^17 → CGBN (L·μ mod n²) ─────────────────────────────── + format_to_cgbn_kernel<<>>(d_prod_cgbn, d_L_b17, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step⑨: m = (L·μ mod n²) mod n ────────────────────────────────────── + // 因 L < n,μ < n,故 L·μ < n²,L·μ mod n² = L·μ,再 mod n 得最终明文 + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + const int threads = 256; + const int blocks = (batch * INV_TPI + threads - 1) / threads; + ou_modp_kernel<<>>(report, d_m_cgbn, d_prod_cgbn, d_n_cgbn, + batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + if (cgbn_error_report_check(report)) + fprintf(stderr, "[paillier_dec_complete] ou_modp_kernel CGBN 错误\n"); + cgbn_error_report_free(report); + } + + // ── Step⑩: CGBN → base-2^17 (m,写入 d_m_out) ────────────────────────── + format_from_cgbn_kernel<<>>(d_m_out, d_m_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + // ── 释放临时设备内存 ───────────────────────────────────────────────────── + cudaFree(d_ct_plain); + cudaFree(d_ct_cgbn); + cudaFree(d_L_cgbn); + cudaFree(d_L_b17); + cudaFree(d_mu_batch); + cudaFree(d_prod_cgbn); + cudaFree(d_m_cgbn); + + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +int main() { + uint64_t Modn[256] = { + 111561, 11897, 65254, 10018, 84184, 4121, 37614, 45117, 28732, + 29823, 114838, 74188, 64256, 85442, 49226, 22722, 59528, 114124, + 37401, 49122, 102504, 72412, 28783, 100394, 121046, 68405, 21607, + 85886, 99460, 55220, 108931, 6714, 123901, 114905, 118498, 85004, + 24806, 113386, 26201, 80098, 119116, 44852, 81005, 59795, 11484, + 116728, 78488, 67258, 120584, 89522, 115707, 48674, 67647, 1877, + 82011, 11196, 28947, 114236, 16867, 28049, 64266, 121733, 46329, + 53338, 110465, 1153, 120566, 9036, 82676, 114714, 121572, 32091, + 103515, 78642, 102261, 121153, 80669, 94055, 30057, 109307, 73935, + 13835, 74325, 55200, 95906, 5145, 114307, 36281, 28556, 27710, + 15408, 23826, 53621, 92615, 350, 104116, 119327, 91218, 80555, + 122208, 21138, 10796, 64727, 90502, 81467, 76724, 27630, 62110, + 37456, 115166, 66131, 83327, 118058, 101425, 116079, 80270, 66310, + 34323, 31475, 79356, 72629, 128248, 115516, 123300, 59893, 112718, + 113186, 39321, 54374, 100387, 41296, 19941, 50310, 17380, 16532, + 99727, 119039, 99958, 65014, 39256, 17782, 58379, 3613, 125771, + 38040, 60277, 39713, 101924, 14157, 65690, 11196, 2427, 68287, + 47464, 84797, 86181, 97818, 81239, 112625, 110466, 66962, 62275, + 76874, 69036, 20298, 16391, 122469, 84269, 87319, 73854, 61208, + 87196, 64610, 64851, 70446, 65674, 1845, 66897, 93817, 102544, + 28867, 84510, 72455, 15626, 38187, 63044, 21375, 101843, 88400, + 61850, 105057, 21554, 12794, 61743, 100541, 27912, 8930, 121065, + 119587, 58351, 67055, 71431, 5672, 69958, 120145, 108909, 94527, + 94235, 117184, 3960, 31237, 24430, 107760, 18040, 39392, 107394, + 104646, 107046, 115355, 129380, 67840, 61436, 67302, 70124, 88363, + 6104, 127851, 117921, 113618, 902, 106694, 28777, 32299, 48675, + 3904, 35777, 48458, 30021, 77854, 31779, 22176, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 49712, 96545, 95680, 129612, 42381, 126558, 4996, 54891, 49903, + 1085, 92362, 67565, 8490, 90336, 104721, 11412, 88468, 17144, + 113144, 112256, 44586, 114095, 43454, 100502, 129722, 38537, 71795, + 61988, 4727, 25298, 11339, 22332, 43802, 62768, 99753, 90230, + 41257, 117280, 61518, 86536, 118921, 40760, 128654, 56100, 112621, + 128100, 40886, 102300, 112819, 104596, 66525, 45541, 102449, 4880, + 67028, 86649, 56788, 100989, 111200, 41395, 18862, 58398, 45639, + 91372, 7724, 117351, 55325, 85810, 32783, 88118, 33084, 105041, + 85813, 60215, 69396, 46478, 34367, 26821, 22778, 17987, 2867, + 29653, 24137, 4486, 87190, 98130, 25654, 67445, 38235, 4885, + 102434, 97969, 96963, 15284, 117818, 73962, 29087, 64482, 94448, + 83602, 26772, 107076, 65936, 40159, 62173, 101570, 67732, 16192, + 17076, 72456, 63017, 114756, 12394, 2159, 127179, 111833, 100888, + 69970, 85720, 99329, 68885, 106268, 112799, 111379, 83009, 33563, + 38777, 53886, 57399, 49067, 95897, 23757, 106219, 30895, 53330, + 44554, 49364, 125896, 68069, 110413, 58091, 109430, 84311, 42699, + 83232, 94778, 42260, 10349, 82115, 113286, 130337, 169, 37893, + 926, 93929, 53931, 91149, 69569, 24269, 112991, 60441, 49438, + 55760, 28404, 112838, 59729, 2987, 55497, 116905, 9223, 78938, + 3084, 3830, 19128, 26821, 77382, 91333, 56237, 93064, 19139, + 68581, 105591, 45380, 60006, 21766, 97161, 130684, 129434, 29297, + 35050, 73454, 79727, 113545, 61388, 118435, 29793, 65311, 120613, + 56284, 26599, 69143, 120840, 20508, 73419, 99429, 6190, 54811, + 69911, 80011, 124881, 31927, 63698, 6523, 104110, 105403, 74351, + 4029, 86578, 25325, 87492, 115337, 101755, 59928, 53336, 106699, + 113368, 75138, 106362, 91726, 69046, 31584, 42609, 35889, 81581, + 100978, 31020, 56316, 4602, 43901, 25722, 56181, 105052, 111796, + 44848, 119076, 27670, 31529, 48743, 34581, 82997, 49852, 39885, + 102542, 67993, 99832, 78469}; + uint64_t R1[256] = { + 125850, 40759, 71657, 107468, 55502, 2623, 101949, 86456, 96860, + 47671, 36778, 48009, 34575, 48990, 29556, 60225, 48278, 107988, + 81193, 45326, 34017, 99906, 6357, 79032, 9259, 97064, 99386, + 37593, 21536, 101253, 58552, 123032, 82224, 110937, 80951, 38093, + 97092, 58863, 111469, 113611, 65470, 60334, 77780, 59294, 106314, + 36681, 51563, 128765, 7020, 74946, 34781, 16340, 32519, 13331, + 53481, 96317, 36652, 48132, 109891, 55273, 53540, 130542, 62285, + 12935, 120, 89414, 46749, 31461, 55333, 104838, 24345, 44395, + 7838, 590, 105503, 35169, 118383, 53798, 33488, 76343, 27984, + 126358, 77080, 60547, 92132, 75301, 3661, 119502, 126208, 39660, + 32852, 100737, 2375, 124455, 3321, 100852, 8138, 129706, 9900, + 41135, 87427, 20220, 27055, 127761, 27886, 24362, 98703, 99562, + 105017, 127786, 79965, 54424, 60755, 39238, 22451, 33060, 39667, + 93604, 63792, 98979, 30520, 100584, 79430, 24603, 96679, 107766, + 108211, 26294, 36587, 101203, 5810, 3400, 68914, 100134, 37375, + 5263, 103586, 100974, 21947, 28008, 118706, 35712, 119326, 13669, + 50624, 25624, 21975, 52741, 108344, 111372, 17508, 8275, 30523, + 55435, 126734, 98681, 49726, 39692, 68397, 93166, 91996, 92587, + 118478, 69402, 124810, 86819, 25315, 62093, 129823, 3163, 112582, + 91399, 39880, 116751, 74966, 48494, 95717, 91858, 98930, 75473, + 1517, 83067, 127014, 97086, 27602, 3392, 58215, 110333, 8766, + 108559, 47083, 116393, 101862, 122911, 23285, 89240, 34295, 116078, + 33505, 104469, 130885, 117195, 11843, 16396, 82096, 98243, 39678, + 24007, 73946, 18685, 6158, 81677, 95553, 60522, 113440, 29404, + 72039, 120808, 35494, 22752, 79555, 114896, 63611, 78608, 125322, + 18038, 127789, 11507, 122079, 30778, 95317, 73874, 124195, 22895, + 114409, 124639, 92400, 127331, 52061, 126224, 5730, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 6530, 128316, 85241, 78089, 44125, 104321, 89348, 78092, 126628, + 25209, 23950, 111858, 46284, 118426, 34408, 52497, 118323, 18182, + 4066, 65455, 99185, 74161, 34036, 79120, 10056, 100447, 121656, + 108328, 28794, 124724, 51228, 22024, 49998, 87909, 116146, 38899, + 91473, 66156, 130063, 70154, 33961, 65086, 70058, 88427, 53894, + 47218, 82342, 93069, 79816, 64412, 40467, 107794, 5891, 47974, + 102570, 86780, 20963, 41782, 91203, 735, 63176, 109925, 24098, + 25942, 112312, 115659, 118407, 21592, 66033, 72266, 99797, 3770, + 34166, 18644, 98585, 85237, 12839, 1276, 72733, 118869, 100014, + 53901, 119833, 95570, 64971, 74765, 111586, 112724, 100897, 78016, + 29024, 8981, 20200, 55894, 129347, 63499, 99987, 96498, 21726, + 95588, 3827, 53586, 70180, 116172, 27352, 90271, 44236, 90262, + 74733, 87994, 113335, 6074, 45626, 75650, 130263, 105461, 54879, + 61529, 67417, 27819, 6223, 109863, 93665, 16699, 79283, 127355, + 71325, 117082, 86035, 19627, 106947, 59040, 93273, 113034, 57994, + 91220, 62548, 117082, 64240, 60903, 72264, 88469, 108683, 47729, + 65893, 109835, 47258, 118821, 125624, 55290, 52757, 119708, 45337, + 15432, 36104, 114440, 94324, 98100, 35869, 106865, 99167, 14239, + 127884, 45415, 116411, 57565, 50087, 34803, 68051, 21734, 33669, + 129689, 5072, 62802, 92055, 83795, 107011, 3981, 5771, 50933, + 112841, 88014, 71962, 32285, 108965, 44261, 6967, 24953, 120162, + 106314, 27069, 75264, 83398, 75587, 108964, 43088, 25635, 7524, + 105113, 71056, 59520, 23835, 25332, 58434, 59962, 108382, 58498, + 8083, 5068, 17102, 60696, 129427, 104426, 25693, 91197, 1907, + 78068, 98160, 104024, 103361, 81350, 55058, 69633, 29378, 83115, + 9700, 243, 106809, 53594, 63755, 122146, 61770, 36653, 73663, + 98577, 119808, 103416, 12753, 98276, 109143, 3003, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cudaMemcpy((void*)Modn,(void*)d_Modn, bytes, cudaMemcpyDeviceToHost); + cout<<"nct n"<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + /* + cudaMemcpy((void*)negModn,(void*)d_negModn, bytes, cudaMemcpyDeviceToHost); + cout<<"ct n'"<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + // ── paillier_dec_complete 测试 ──────────────────────────────────────────── + { + const int BATCH = 200000; + const int PAILLIER_LAMBDA_TAU = 2048; // λ ≈ 2047 bit,对齐到 64 + const size_t ct_bytes = (size_t)BATCH * ARR_LEN * sizeof(uint64_t); + + // 测试明文:m[i] = i + 1(均 < n) + const uint64_t test_m[BATCH] = {1, 2, 3, 4, 5, 6, 7, 8}; + + // 构造标准域密文:c = 1 + m·n(无随机项,base-2^17 大整数乘法) + std::vector h_ct_plain((size_t)BATCH * ARR_LEN, 0); + for (int i = 0; i < BATCH; i++) { + uint64_t *limbs = h_ct_plain.data() + (size_t)i * ARR_LEN; + uint64_t carry = 0; + for (int j = 0; j < ARR_LEN; j++) { + uint64_t prod = test_m[i] * paillier_n[j] + carry; + limbs[j] = prod & 0x1FFFFULL; + carry = prod >> 17; + } + uint64_t ac = 1; + for (int j = 0; j < ARR_LEN && ac; j++) { + uint64_t s = limbs[j] + ac; + limbs[j] = s & 0x1FFFFULL; + ac = s >> 17; + } + } + + uint64_t *d_ct_tilde_test = nullptr; + uint64_t *d_m_out_test = nullptr; + CUDA_CHECK(cudaMalloc(&d_ct_tilde_test, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_out_test, ct_bytes)); + + // ── 准备密钥设备缓冲区(λ、μ、n)──────────────────────────────────── + const int exp_limbs_key = PAILLIER_LAMBDA_TAU / 64; // 32 + + // 主机端:λ base-2^17 → base-2^64,复制 BATCH 份 + std::vector lambda_batch_key((size_t)BATCH * exp_limbs_key); + { + std::vector lam64(exp_limbs_key); + bn17_to_bn64(paillier_lambda, lam64.data(), exp_limbs_key); + for (int i = 0; i < BATCH; i++) + memcpy(lambda_batch_key.data() + (size_t)i * exp_limbs_key, + lam64.data(), exp_limbs_key * sizeof(uint64_t)); + } + // 主机端:n base-2^17 → CGBN + inv_bn_mem_t h_n_cgbn_key; + bn17_to_cgbn_mem(paillier_n, &h_n_cgbn_key); + + uint64_t *d_lambda_exp_key = nullptr; + uint64_t *d_mu_key = nullptr; + inv_bn_mem_t *d_n_cgbn_key = nullptr; + CUDA_CHECK(cudaMalloc(&d_lambda_exp_key, + (size_t)BATCH * exp_limbs_key * sizeof(uint64_t))); + CUDA_CHECK(cudaMalloc(&d_mu_key, ARR_LEN * sizeof(uint64_t))); + CUDA_CHECK(cudaMalloc(&d_n_cgbn_key, sizeof(inv_bn_mem_t))); + + // ── H2D 计时:密文 + λ + μ + n ─────────────────────────────────────── + cudaEvent_t ev_h2d_s, ev_h2d_e; + CUDA_CHECK(cudaEventCreate(&ev_h2d_s)); + CUDA_CHECK(cudaEventCreate(&ev_h2d_e)); + + CUDA_CHECK(cudaEventRecord(ev_h2d_s)); + CUDA_CHECK(cudaMemcpy(d_ct_tilde_test, h_ct_plain.data(), ct_bytes, + cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_lambda_exp_key, lambda_batch_key.data(), + (size_t)BATCH * exp_limbs_key * sizeof(uint64_t), + cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_mu_key, paillier_mu, ARR_LEN * sizeof(uint64_t), + cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_n_cgbn_key, &h_n_cgbn_key, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev_h2d_e)); + CUDA_CHECK(cudaEventSynchronize(ev_h2d_e)); + + float h2d_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&h2d_ms, ev_h2d_s, ev_h2d_e)); + CUDA_CHECK(cudaEventDestroy(ev_h2d_s)); + CUDA_CHECK(cudaEventDestroy(ev_h2d_e)); + + // c_tilde = c · R mod n²(转 FMLM 域,不计入 H2D 时间) + XYfixWarpROneVector<<>>( + d_ct_tilde_test, d_ctR, d_nctR, d_negModn, d_con_NegModn_shoup, d_Modn, + d_con_Modn_shoup, (uint64_t)ARR_LEN, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── GPU 解密计时 ─────────────────────────────────────────────────────── + float dec_ms = paillier_dec_complete( + d_ct_tilde_test, d_m_out_test, BATCH, PAILLIER_LAMBDA_TAU, + d_lambda_exp_key, d_mu_key, d_n_cgbn_key, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, d_ctR, d_nctR); + + // ── D2H 计时 ────────────────────────────────────────────────────────── + std::vector h_m_out((size_t)BATCH * ARR_LEN, 0); + + cudaEvent_t ev_d2h_s, ev_d2h_e; + CUDA_CHECK(cudaEventCreate(&ev_d2h_s)); + CUDA_CHECK(cudaEventCreate(&ev_d2h_e)); + + CUDA_CHECK(cudaEventRecord(ev_d2h_s)); + CUDA_CHECK(cudaMemcpy(h_m_out.data(), d_m_out_test, ct_bytes, + cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev_d2h_e)); + CUDA_CHECK(cudaEventSynchronize(ev_d2h_e)); + + float d2h_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&d2h_ms, ev_d2h_s, ev_d2h_e)); + CUDA_CHECK(cudaEventDestroy(ev_d2h_s)); + CUDA_CHECK(cudaEventDestroy(ev_d2h_e)); + + // ── 打印时间(us)──────────────────────────────────────────────────── + printf("========== paillier_dec_complete timing (batch=%d) ==========\n", + BATCH); + printf(" H2D (ct+key): %10.2f us\n", h2d_ms * 1000.f); + printf(" GPU decrypt : %10.2f us\n", dec_ms * 1000.f); + printf(" D2H (result): %10.2f us\n", d2h_ms * 1000.f); + printf(" Total : %10.2f us\n", (h2d_ms + dec_ms + d2h_ms) * 1000.f); + printf(" Avg per dec : %10.2f us (GPU only / batch)\n", + dec_ms * 1000.f / BATCH); + printf("=============================================================\n"); + /* +// ── 写入验证文件(供 Python 脚本验证正确性)───────────────────────── +FILE *fp = fopen("paillier_dec_complete_test.txt", "w"); +if (fp) { + fprintf(fp, "# paillier_dec_complete test output\n"); + fprintf(fp, "N:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", (unsigned long +long)paillier_n[j]); fprintf(fp, "\n"); fprintf(fp, "LAMBDA:"); for (int j = 0; +j < ARR_LEN; j++) fprintf(fp, " %llu", (unsigned long long)paillier_lambda[j]); + fprintf(fp, "\n"); + fprintf(fp, "MU:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", (unsigned long +long)paillier_mu[j]); fprintf(fp, "\n"); for (int i = 0; i < BATCH; i++) { + fprintf(fp, "M: %llu", (unsigned long long)test_m[i]); + for (int j = 1; j < 32; j++) fprintf(fp, " 0"); + fprintf(fp, "\n"); + const uint64_t *ct = h_ct_plain.data() + (size_t)i * ARR_LEN; + fprintf(fp, "CT:"); + for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " %llu", (unsigned long +long)ct[j]); fprintf(fp, "\n"); const uint64_t *md = h_m_out.data() + (size_t)i +* ARR_LEN; fprintf(fp, "MDEC:"); for (int j = 0; j < ARR_LEN; j++) fprintf(fp, " +%llu", (unsigned long long)md[j]); fprintf(fp, "\n"); + } + fclose(fp); + printf("验证文件已写入: paillier_dec_complete_test.txt\n"); +} + */ + cudaFree(d_ct_tilde_test); + cudaFree(d_m_out_test); + cudaFree(d_lambda_exp_key); + cudaFree(d_mu_key); + cudaFree(d_n_cgbn_key); + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_enc.cu b/heu/library/algorithms/paillier_new/paillier_enc.cu new file mode 100644 index 0000000..dd8e2fa --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_enc.cu @@ -0,0 +1,18351 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 16 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + /* + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 1000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + n²))──────────────────────────────── srand((unsigned int)time(NULL)); for + (int p = 0; p < NUM_DEC; p++) { uint64_t *c = h_ct + (size_t)p * ARR_LEN; for + (int j = 0; j < ARR_LEN; j++) { if (j < n2_top_idx) { c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & mask17; } + else if (j == n2_top_idx) { c[j] = (uint64_t)(rand() % (int)n2_top_val); } + else { c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = ((uint64_t)(uint32_t)rand() << 32) | + (uint64_t)(uint32_t)rand(); for (int p = 0; p < NUM_DEC; p++) memcpy(h_exp + + (size_t)p * EXP_U64_LIMBS, lambda_limbs, EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec ) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, + ARR_LEN, BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", + r, round_ms[r], round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + */ + // ============================================================ + // [paillier_encryptionhs 性能测试] + // Enc(m_i, r_i) = (1 + m_i·n) · h_s^{r_i} mod n²(FMLM域) + // Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域) + // Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) + // Step③ XYfixWarpROneVector:(1+m·n) → FMLM域 + // Step④ XYfixWarpVector:(1̃+m·n) · h̃_s^{r_i} → c̃(FMLM域) + // 结果写入 encryption_results.txt,供 verify_encryptionhs.py 验证正确性。 + // ============================================================ + const int NUM_ENC = 1000; + const int kWarmup = 3; + const int kRounds = 2; + // ── 参数定义 ───────────────────────────────────────────────────────────── + const int TAU_R = 1024; // r ← Z_{2^{k/2}},k=2048 + const int EXP_R_U64 = TAU_R / 64; // = 16 + const size_t ct_bytes = (size_t)NUM_ENC * ARR_LEN * sizeof(uint64_t); + const size_t m_bytes = (size_t)NUM_ENC * ARR_LEN * sizeof(uint64_t); + const size_t r_bytes = (size_t)NUM_ENC * EXP_R_U64 * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + + printf("\n============================================================\n"); + printf(" [paillier_encryptionhs] 批量加密性能测试\n"); + printf(" 批大小=%d TAU_R=%d bit WINDOW=%d bit TABLE_SIZE=%d 项\n", + NUM_ENC, TAU_R, WINDOW_BITS, TABLE_SIZE); + printf(" 窗口数=%d(每窗口 %d 次平方 + 1 次乘法)\n", + (TAU_R + WINDOW_BITS - 1) / WINDOW_BITS, WINDOW_BITS); + printf("============================================================\n"); + + // ── 1. CPU 建立 hs 预计算表 ────────────────────────────────────────────── + uint64_t h_hs[ARR_LEN] = {}; + uint64_t *h_table = (uint64_t *)malloc(tbl_bytes); + if (!h_table) { + fprintf(stderr, "[Enc] h_table malloc 失败\n"); + goto enc_done; + } + generate_hs_table(N2_arr, h_hs, h_table); + + { + // ── 2. 分配 pinned 内存 ─────────────────────────────────────────────── + uint64_t *h_m = NULL; // 明文 m(标准域,base-2^BASE_BITS) + uint64_t *h_r = NULL; // 随机指数 r_i(1024-bit,各独立) + uint64_t *h_ct = NULL; // D2H 后密文结果 + CUDA_CHECK(cudaMallocHost(&h_m, m_bytes)); + CUDA_CHECK(cudaMallocHost(&h_r, r_bytes)); + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + + // ── 3. 随机生成输入数据 ─────────────────────────────────────────────── + srand((unsigned int)time(NULL)); + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + // m < n:使用 N_WORDS-1=120 个随机 limb(limb 120..255 填 0),保证 m < n + for (int p = 0; p < NUM_ENC; p++) { + uint64_t *mp = h_m + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < N_WORDS - 1) // 0..119:随机 limb in [0, BASE) + mp[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + else + mp[j] = 0ULL; // 120..255:填 0,确保 m < n + } + } + // r_i:各独立的 1024-bit 随机数 + for (int p = 0; p < NUM_ENC; p++) { + uint64_t *rp = h_r + (size_t)p * EXP_R_U64; + for (int j = 0; j < EXP_R_U64; j++) + rp[j] = ((uint64_t)(uint32_t)rand() << 32) | (uint64_t)(uint32_t)rand(); + } + printf("[Enc] 随机数据生成完成(%d 个明文,各 %d-bit r_i)\n", NUM_ENC, + TAU_R); + + // ── 4. 分配 GPU 内存 & H2D ──────────────────────────────────────────── + uint64_t *d_table = NULL; + uint64_t *d_m = NULL; + uint64_t *d_r = NULL; + uint64_t *d_ct_out = NULL; + CUDA_CHECK(cudaMalloc(&d_table, tbl_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, m_bytes)); + CUDA_CHECK(cudaMalloc(&d_r, r_bytes)); + CUDA_CHECK(cudaMalloc(&d_ct_out, ct_bytes)); + { + struct timespec t0h, t1h; + clock_gettime(CLOCK_MONOTONIC, &t0h); + CUDA_CHECK( + cudaMemcpy(d_table, h_table, tbl_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_m, h_m, m_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_r, h_r, r_bytes, cudaMemcpyHostToDevice)); + clock_gettime(CLOCK_MONOTONIC, &t1h); + double h2d_ms = (double)(t1h.tv_sec - t0h.tv_sec) * 1e3 + + (double)(t1h.tv_nsec - t0h.tv_nsec) * 1e-6; + printf("[Enc] H2D: %.3f ms 表 %.2f MB + 明文 %.2f MB + 指数 %.2f MB\n", + h2d_ms, tbl_bytes / 1048576.0, m_bytes / 1048576.0, + r_bytes / 1048576.0); + } + + // ── 5. 预热 ─────────────────────────────────────────────────────────── + printf("[Enc] 预热 %d 次...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_encryptionhs( + d_m, d_ct_out, d_N, d_N2, d_table, d_r, TAU_R, NUM_ENC, d_ctR, d_nctR, + d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaDeviceSynchronize()); + } + printf("[Enc] 预热完成\n\n"); + + // ── 6. cudaEvent 计时(d_m/d_r 只读,d_ct_out 每轮写入,无需重置)────── + printf("[Enc] 开始计时(%d 轮 × %d 个明文)...\n", kRounds, NUM_ENC); + float rms[5] = {}; + for (int rnd = 0; rnd < kRounds; rnd++) { + rms[rnd] = paillier_encryptionhs( + d_m, d_ct_out, d_N, d_N2, d_table, d_r, TAU_R, NUM_ENC, d_ctR, d_nctR, + d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + } + printf("[Enc] 计时完成\n\n"); + + // ── 7. D2H 最后一轮密文结果 ────────────────────────────────────────── + { + struct timespec t0d, t1d; + clock_gettime(CLOCK_MONOTONIC, &t0d); + CUDA_CHECK(cudaMemcpy(h_ct, d_ct_out, ct_bytes, cudaMemcpyDeviceToHost)); + clock_gettime(CLOCK_MONOTONIC, &t1d); + double d2h_ms = (double)(t1d.tv_sec - t0d.tv_sec) * 1e3 + + (double)(t1d.tv_nsec - t0d.tv_nsec) * 1e-6; + printf("[Enc] D2H: %.3f ms 密文 %.2f MB\n", d2h_ms, + ct_bytes / 1048576.0); + } + + // ── 8. 写结果到 encryption_results.txt ────────────────────────────── + // 格式(供 verify_encryptionhs.py 验证): + // 行1: NUM_ENC ARR_LEN BASE_BITS TAU_R + // 行2: n² limbs(ARR_LEN 个 uint64,base-2^BASE_BITS,小端序) + // 行3: n limbs(ARR_LEN 个 uint64) + // 行4: r_param = table[0](FMLM 恒等元) + // 行5: hs limbs(ARR_LEN 个 uint64,标准域) + // 每个用例 3 行:m limbs / r_i limbs / c̃ limbs + { + FILE *fout = fopen("encryption_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Enc] 无法创建 encryption_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_ENC, ARR_LEN, BASE_BITS, TAU_R); + // n² + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // n + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N_arr[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // r_param = table[0](FMLM 恒等元) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)h_table[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // hs(底数,标准域) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)h_hs[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // 每个用例 3 行 + for (int p = 0; p < NUM_ENC; p++) { + const uint64_t *mp = h_m + (size_t)p * ARR_LEN; + const uint64_t *rp = h_r + (size_t)p * EXP_R_U64; + const uint64_t *ctp = h_ct + (size_t)p * ARR_LEN; + // m(明文,标准域) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)mp[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + // r_i(随机指数,16 个 uint64) + for (int j = 0; j < EXP_R_U64; j++) + fprintf(fout, "%llu%c", (unsigned long long)rp[j], + j + 1 < EXP_R_U64 ? ' ' : '\n'); + // c̃(密文,FMLM域) + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)ctp[j], + j + 1 < ARR_LEN ? ' ' : '\n'); + } + fclose(fout); + printf("[Enc] 结果已写入 encryption_results.txt\n"); + } + } + + // ── 9. 性能统计 ─────────────────────────────────────────────────────── + float rmin = rms[0], rmax = rms[0], rsum = 0.f; + for (int rnd = 0; rnd < kRounds; rnd++) { + rsum += rms[rnd]; + if (rms[rnd] < rmin) rmin = rms[rnd]; + if (rms[rnd] > rmax) rmax = rms[rnd]; + } + float ravg = rsum / kRounds; + printf("\n============================================================\n"); + printf(" paillier_encryptionhs 性能报告\n"); + printf(" 批大小=%d TAU_R=%d bit WINDOW=%d bit\n", NUM_ENC, TAU_R, + WINDOW_BITS); + printf("============================================================\n"); + printf("[计时结果(cudaEvent,仅 GPU 核函数,不含传输)]\n"); + printf("------------------------------------------------------------\n"); + for (int rnd = 0; rnd < kRounds; rnd++) + printf(" 轮 %d : %8.3f ms (%.3f us/次)\n", rnd, rms[rnd], + rms[rnd] * 1e3f / NUM_ENC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %8.3f ms\n", ravg); + printf(" 最快轮 : %8.3f ms\n", rmin); + printf(" 最慢轮 : %8.3f ms\n", rmax); + printf(" 平均单次 Encrypt : %8.3f us\n", ravg * 1e3f / NUM_ENC); + printf(" (最快轮) 单次 : %8.3f us\n", rmin * 1e3f / NUM_ENC); + printf("============================================================\n"); + + // ── 10. 释放 Encryption 相关内存 ────────────────────────────────────── + cudaFreeHost(h_m); + cudaFreeHost(h_r); + cudaFreeHost(h_ct); + cudaFree(d_table); + cudaFree(d_m); + cudaFree(d_r); + cudaFree(d_ct_out); + } + free(h_table); + printf("[Enc] 测试完成\n"); +enc_done:; + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_encryptionhs 测试完成。\n"); + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_mulp.cu b/heu/library/algorithms/paillier_new/paillier_mulp.cu new file mode 100644 index 0000000..0b22c78 --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_mulp.cu @@ -0,0 +1,16970 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +/// #include +#include + +#include +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 200000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 1024 // 指数 bit 数 +#define EXP_U64_LIMBS (TAU / 64) // = 16,压缩格式每个指数占 16 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 10000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0, 0, 0, 0, 0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 (1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + /* + // ── 锁页内存分配 ───────────────────────────────────────────── + const size_t ct_bytes = (size_t)TEST_BATCH * ARR_LEN * + sizeof(uint64_t); const size_t exp_bytes = (size_t)TEST_BATCH * EXP_U64_LIMBS + * sizeof(uint64_t); + + uint64_t *h_bases, *h_exps, *h_results; + CUDA_CHECK(cudaMallocHost(&h_bases, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exps, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_results, ct_bytes)); + memset(h_results, 0, ct_bytes); + + // ── 生成测试数据 ───────────────────────────────────────────── + // + // 底数 x:直接随机生成 < n²(天然是蒙哥马利域元素) + // x 代表普通域值 a = x · R⁻¹ mod n² + // Python 会用 R⁻¹ 还原 a,再验证 a^m · R + // + // 指数 m:1024-bit 随机数(压缩为 16 个 uint64_t) + // + srand((unsigned int)time(NULL)); + const uint64_t mask17 = BASE - 1; + + for (int p = 0; p < TEST_BATCH; p++) { + // ── 底数(蒙哥马利域,直接随机生成 < n²)────────────────── + uint64_t *x = h_bases + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < N2_TOP_IDX) { + // limb 0~239:随机 [0, BASE) + x[j] = ((uint64_t)rand() ^ ((uint64_t)rand() << 15)) & mask17; + } else if (j == N2_TOP_IDX) { + // limb 240:随机 [0, N2_TOP_VAL),保证 x < n² + x[j] = (uint64_t)rand() % N2_TOP_VAL; + } else { + // limb 241~255:0 + x[j] = 0ULL; + } + } + + // ── 指数(1024-bit,压缩为 16 个 uint64_t)───────────────── + uint64_t *e = h_exps + (size_t)p * EXP_U64_LIMBS; + for (int k = 0; k < EXP_U64_LIMBS; k++) { + uint64_t v = 0; + v |= (uint64_t)(rand() & 0x7FFF); + v |= (uint64_t)(rand() & 0x7FFF) << 15; + v |= (uint64_t)(rand() & 0x7FFF) << 30; + v |= (uint64_t)(rand() & 0x7FFF) << 45; + v |= (uint64_t)(rand() & 0x000F) << 60; + e[k] = v; + } + // 确保最高 bit 为 1(使指数恰好 1024-bit) + e[EXP_U64_LIMBS - 1] |= (1ULL << 63); + } + + printf("已生成 %d 个测试用例(底数随机 < n²,指数 1024-bit)\n\n", + TEST_BATCH); + + // ── GPU 内存分配 ───────────────────────────────────────────── + uint64_t *d_bases, *d_exps, *d_output; + CUDA_CHECK(cudaMalloc(&d_bases, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exps, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_output, ct_bytes)); + + // ── H2D ────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(d_bases, h_bases, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exps, h_exps, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 内核启动 ───────────────────────────────────────────────── + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + const int n_blocks = (TEST_BATCH + WARP_PER_BLK - 1) / WARP_PER_BLK; + + { + cudaDeviceProp prop; + cudaGetDeviceProperties(&prop, 0); + printf("GPU: %s\n", prop.name); + printf("共享内存申请: %zu KB / 最大 %zu KB\n\n", + smem_size / 1024, prop.sharedMemPerBlock / 1024); + if (smem_size > prop.sharedMemPerBlock) { + fprintf(stderr, "[ERROR] 共享内存超限\n"); return -1; + } + } + + printf("启动 FMLE_mod2_Kernel: %d blocks × %d threads\n", + n_blocks, WARP_PER_BLK * 32); + + FMLE_mod2_Kernel<<>>( + d_bases, + d_exps, + TAU, + d_output, + TEST_BATCH, + d_r_0, + // ★ 无 d_r1_ct / d_r1_nct(Step8 跳过) + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + 256, MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, + d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H ────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_results, d_output, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 写 txt ─────────────────────────────────────────────────── + write_to_txt("mulplain_mod2_results.txt", + TEST_BATCH, h_bases, h_exps, h_results, Modn); + + // ── 清理 ───────────────────────────────────────────────────── + CUDA_CHECK(cudaFreeHost(h_bases)); + CUDA_CHECK(cudaFreeHost(h_exps)); + CUDA_CHECK(cudaFreeHost(h_results)); + + CUDA_CHECK(cudaFree(d_bases)); + CUDA_CHECK(cudaFree(d_exps)); + CUDA_CHECK(cudaFree(d_output)); + CUDA_CHECK(cudaFree(d_con_twiddle)); + CUDA_CHECK(cudaFree(d_con_twiddle_shoup)); + CUDA_CHECK(cudaFree(d_con_twiddle_NCT)); + CUDA_CHECK(cudaFree(d_con_twiddle_NCT_shoup)); + CUDA_CHECK(cudaFree(d_con_InvTwiddle)); + CUDA_CHECK(cudaFree(d_con_InvTwiddle_shoup)); + CUDA_CHECK(cudaFree(d_con_ICTTwiddle)); + CUDA_CHECK(cudaFree(d_con_ICTTwiddle_shoup)); + CUDA_CHECK(cudaFree(d_Modn)); + CUDA_CHECK(cudaFree(d_con_Modn_shoup)); + CUDA_CHECK(cudaFree(d_negModn)); + CUDA_CHECK(cudaFree(d_con_NegModn_shoup)); + CUDA_CHECK(cudaFree(d_sample)); + CUDA_CHECK(cudaFree(d_r_0)); + + printf("测试完成,请运行 Python 验证脚本。\n"); + return 0; + */ + // ── 锁页内存 ───────────────────────────────────────────────── + + const size_t ct_bytes = (size_t)TEST_BATCH * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = + (size_t)TEST_BATCH * EXP_U64_LIMBS * sizeof(uint64_t); + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + uint64_t *h_bases, *h_exps, *h_results; + CUDA_CHECK(cudaMallocHost(&h_bases, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exps, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_results, ct_bytes)); + memset(h_results, 0, ct_bytes); + + // ── 生成测试数据 ───────────────────────────────────────────── + srand((unsigned int)time(NULL)); + const uint64_t mask17 = BASE - 1; + + for (int p = 0; p < TEST_BATCH; p++) { + uint64_t *x = h_bases + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < N2_TOP_IDX) + x[j] = ((uint64_t)rand() ^ ((uint64_t)rand() << 15)) & mask17; + else if (j == N2_TOP_IDX) + x[j] = (uint64_t)rand() % N2_TOP_VAL; + else + x[j] = 0ULL; + } + uint64_t *e = h_exps + (size_t)p * EXP_U64_LIMBS; + for (int k = 0; k < EXP_U64_LIMBS; k++) { + uint64_t v = 0; + v |= (uint64_t)(rand() & 0x7FFF); + v |= (uint64_t)(rand() & 0x7FFF) << 15; + v |= (uint64_t)(rand() & 0x7FFF) << 30; + v |= (uint64_t)(rand() & 0x7FFF) << 45; + v |= (uint64_t)(rand() & 0x000F) << 60; + e[k] = v; + } + e[EXP_U64_LIMBS - 1] |= (1ULL << 63); + } + printf("已生成 %d 个测试用例\n\n", TEST_BATCH); + + // ── GPU 数据缓冲区分配 ─────────────────────────────────────── + uint64_t *d_bases, *d_exps, *d_output; + CUDA_CHECK(cudaMalloc(&d_bases, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exps, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_output, ct_bytes)); + + // ============================================================ + // H2D:数据上传(计时) + // ============================================================ + cudaEvent_t ev_h2d_s, ev_h2d_e; + CUDA_CHECK(cudaEventCreate(&ev_h2d_s)); + CUDA_CHECK(cudaEventCreate(&ev_h2d_e)); + + CUDA_CHECK(cudaEventRecord(ev_h2d_s)); + CUDA_CHECK(cudaMemcpy(d_bases, h_bases, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exps, h_exps, exp_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev_h2d_e)); + CUDA_CHECK(cudaEventSynchronize(ev_h2d_e)); + + float h2d_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&h2d_ms, ev_h2d_s, ev_h2d_e)); + cudaEventDestroy(ev_h2d_s); + cudaEventDestroy(ev_h2d_e); + + // ============================================================ + // Warm-up:预热 GPU(消除 JIT/升频/Cache 冷启动影响) + // ============================================================ + printf("[Warm-up] 开始...\n"); + { + MulPlainTiming dummy = paillier_mulplain( + d_bases, d_exps, d_output, TEST_BATCH, TAU, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + (void)dummy; + CUDA_CHECK(cudaDeviceSynchronize()); + } + printf("[Warm-up] 完成\n\n"); + + // ============================================================ + // 正式调用 paillier_mulplain(只做计算,无 H2D/D2H) + // ============================================================ + printf("[正式运行] paillier_mulplain(%d 流)\n", NUM_STREAMS); + MulPlainTiming kt = paillier_mulplain( + d_bases, d_exps, d_output, TEST_BATCH, TAU, d_r_0, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + + // ============================================================ + // D2H:结果下载(计时) + // ============================================================ + cudaEvent_t ev_d2h_s, ev_d2h_e; + CUDA_CHECK(cudaEventCreate(&ev_d2h_s)); + CUDA_CHECK(cudaEventCreate(&ev_d2h_e)); + + CUDA_CHECK(cudaEventRecord(ev_d2h_s)); + CUDA_CHECK(cudaMemcpy(h_results, d_output, ct_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev_d2h_e)); + CUDA_CHECK(cudaEventSynchronize(ev_d2h_e)); + + float d2h_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&d2h_ms, ev_d2h_s, ev_d2h_e)); + cudaEventDestroy(ev_d2h_s); + cudaEventDestroy(ev_d2h_e); + + // ── 打印计时 ────────────────────────────────────────────────── + float total_ms = h2d_ms + kt.pipeline_ms + d2h_ms; + printf("\n======================================================\n"); + printf(" paillier_mulplain 测试结果\n"); + printf("======================================================\n"); + printf(" 测试用例数 : %d\n", TEST_BATCH); + printf(" 指数 bit 数 : %d\n", TAU); + printf(" 流数量 : %d\n", NUM_STREAMS); + printf(" 每流用例数 : %d\n", TEST_BATCH / NUM_STREAMS); + printf("------------------------------------------------------\n"); + printf(" H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" 内核计算 : %8.4f ms (各流最大值)\n", kt.kernel_ms); + printf(" 流水线总时间 : %8.4f ms (含流调度开销)\n", kt.pipeline_ms); + printf(" D2H 下载 : %8.4f ms\n", d2h_ms); + printf("------------------------------------------------------\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次内核 : %8.4f us/次\n", kt.kernel_ms * 1000.f / TEST_BATCH); + printf(" 吞吐量(内核): %.0f ops/s\n", + TEST_BATCH / (kt.kernel_ms / 1000.f)); + printf("======================================================\n\n"); + + // ── 写 txt 供 Python 验证 ───────────────────────────────────── + // write_to_txt("mulplain_mod2_results.txt", + // TEST_BATCH, h_bases, h_exps, h_results, Modn); + + // ── 清理 ───────────────────────────────────────────────────── + CUDA_CHECK(cudaFreeHost(h_bases)); + CUDA_CHECK(cudaFreeHost(h_exps)); + CUDA_CHECK(cudaFreeHost(h_results)); + + CUDA_CHECK(cudaFree(d_bases)); + CUDA_CHECK(cudaFree(d_exps)); + CUDA_CHECK(cudaFree(d_output)); + CUDA_CHECK(cudaFree(d_con_twiddle)); + CUDA_CHECK(cudaFree(d_con_twiddle_shoup)); + CUDA_CHECK(cudaFree(d_con_twiddle_NCT)); + CUDA_CHECK(cudaFree(d_con_twiddle_NCT_shoup)); + CUDA_CHECK(cudaFree(d_con_InvTwiddle)); + CUDA_CHECK(cudaFree(d_con_InvTwiddle_shoup)); + CUDA_CHECK(cudaFree(d_con_ICTTwiddle)); + CUDA_CHECK(cudaFree(d_con_ICTTwiddle_shoup)); + CUDA_CHECK(cudaFree(d_Modn)); + CUDA_CHECK(cudaFree(d_con_Modn_shoup)); + CUDA_CHECK(cudaFree(d_negModn)); + CUDA_CHECK(cudaFree(d_con_NegModn_shoup)); + CUDA_CHECK(cudaFree(d_sample)); + CUDA_CHECK(cudaFree(d_r_0)); + + printf("测试完成,请运行:python verify_fmle_mod2.py\n"); + return 0; + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_randc.cu b/heu/library/algorithms/paillier_new/paillier_randc.cu new file mode 100644 index 0000000..b4c341a --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_randc.cu @@ -0,0 +1,18235 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 12 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + /* + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 1000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + n²))──────────────────────────────── srand((unsigned int)time(NULL)); for + (int p = 0; p < NUM_DEC; p++) { uint64_t *c = h_ct + (size_t)p * ARR_LEN; for + (int j = 0; j < ARR_LEN; j++) { if (j < n2_top_idx) { c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & mask17; } + else if (j == n2_top_idx) { c[j] = (uint64_t)(rand() % (int)n2_top_val); } + else { c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = ((uint64_t)(uint32_t)rand() << 32) | + (uint64_t)(uint32_t)rand(); for (int p = 0; p < NUM_DEC; p++) memcpy(h_exp + + (size_t)p * EXP_U64_LIMBS, lambda_limbs, EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec ) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, + ARR_LEN, BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", + r, round_ms[r], round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + */ + // ============================================================ + // [paillier_randomize2 性能测试] + // Randomize(c̃₁) = FMLM(c̃₁, h̃_s^{r_i}) mod n²(FMLM域) + // Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域) + // Step② XYfixWarpVector:c̃₁ · h̃_s^{r_i} → c̃_out + // 结果写入 randomize2_results.txt,供 verify_randomize2.py 验证正确性。 + // ============================================================ + const int NUM_DEC = 100000; // 复用变量名,保持后续代码兼容 + const int kWarmup = 3; + const int kRounds = 2; + // ── 参数定义 ───────────────────────────────────────────────────────────── + const int NUM_RND = NUM_DEC; // 1000 个密文 + const int TAU_R = 1024; // r ← Z_{2^{k/2}},k=2048 + const int EXP_R_U64 = TAU_R / 64; // = 16 + const size_t ct_bytes = (size_t)NUM_RND * ARR_LEN * sizeof(uint64_t); + const size_t r_bytes = (size_t)NUM_RND * EXP_R_U64 * sizeof(uint64_t); + const size_t tbl_bytes = (size_t)TABLE_SIZE * ARR_LEN * sizeof(uint64_t); + + printf("\n============================================================\n"); + printf(" [paillier_randomize2] 批量密文随机化性能测试\n"); + printf(" 批大小=%d TAU_R=%d bit WINDOW=%d bit TABLE_SIZE=%d 项\n", + NUM_RND, TAU_R, WINDOW_BITS, TABLE_SIZE); + printf(" 窗口数=%d(每窗口 %d 次平方 + 1 次乘法)\n", + (TAU_R + WINDOW_BITS - 1) / WINDOW_BITS, WINDOW_BITS); + printf("============================================================\n"); + + // ── 1. CPU 建立 hs 预计算表 ────────────────────────────────────────────── + uint64_t h_hs[ARR_LEN] = {}; + uint64_t *h_table = (uint64_t *)malloc(tbl_bytes); + if (!h_table) { + fprintf(stderr, "[Rand2] h_table malloc 失败\n"); + goto rand2_done; + } + generate_hs_table(N2_arr, h_hs, h_table); + + { + // ── 2. 分配 pinned 内存 ─────────────────────────────────────────────── + uint64_t *h_c1 = NULL; // 原始输入密文 c̃₁(FMLM域,[0, n²)) + uint64_t *h_r = NULL; // 随机指数 r_i(1024-bit,各独立) + uint64_t *h_c_out = NULL; // D2H 后结果 + CUDA_CHECK(cudaMallocHost(&h_c1, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_r, r_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c_out, ct_bytes)); + + // ── 3. 随机生成输入数据 ─────────────────────────────────────────────── + srand((unsigned int)time(NULL)); + // c̃₁:随机合法 FMLM 域密文,限制 < n²(利用 n² 最高非零 limb 下界) + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + for (int p = 0; p < NUM_RND; p++) { + uint64_t *c = h_c1 + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) + c[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + else if (j == n2_top_idx) + c[j] = (uint64_t)(rand() % (int)n2_top_val); + else + c[j] = 0ULL; + } + } + // r_i:各独立的 1024-bit 随机数 + for (int p = 0; p < NUM_RND; p++) { + uint64_t *rp = h_r + (size_t)p * EXP_R_U64; + for (int j = 0; j < EXP_R_U64; j++) + rp[j] = ((uint64_t)(uint32_t)rand() << 32) | (uint64_t)(uint32_t)rand(); + } + printf("[Rand2] 随机数据生成完成(%d 个密文,各 %d-bit r_i)\n", NUM_RND, + TAU_R); + + // ── 4. 分配 GPU 内存 & H2D ──────────────────────────────────────────── + uint64_t *d_table = NULL; + uint64_t *d_c1 = NULL; + uint64_t *d_r = NULL; + CUDA_CHECK(cudaMalloc(&d_table, tbl_bytes)); + CUDA_CHECK(cudaMalloc(&d_c1, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_r, r_bytes)); + { + struct timespec t0h, t1h; + clock_gettime(CLOCK_MONOTONIC, &t0h); + CUDA_CHECK( + cudaMemcpy(d_table, h_table, tbl_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_r, h_r, r_bytes, cudaMemcpyHostToDevice)); + clock_gettime(CLOCK_MONOTONIC, &t1h); + double h2d_ms = (double)(t1h.tv_sec - t0h.tv_sec) * 1e3 + + (double)(t1h.tv_nsec - t0h.tv_nsec) * 1e-6; + printf("[Rand2] H2D: %.3f ms 表 %.2f MB + 密文 %.2f MB + 指数 %.2f MB\n", + h2d_ms, tbl_bytes / 1048576.0, ct_bytes / 1048576.0, + r_bytes / 1048576.0); + } + + // ── 5. 预热 ─────────────────────────────────────────────────────────── + printf("[Rand2] 预热 %d 次...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + // 每次预热前恢复 d_c1(paillier_randomize2 原地修改) + CUDA_CHECK(cudaMemcpy(d_c1, h_c1, ct_bytes, cudaMemcpyHostToDevice)); + paillier_randomize2(d_c1, d_table, d_r, TAU_R, NUM_RND, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaDeviceSynchronize()); + } + printf("[Rand2] 预热完成\n\n"); + + // ── 6. cudaEvent 计时(每轮前重置 d_c1,确保输入一致)───────────────── + printf("[Rand2] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_RND); + float rms[5] = {}; + for (int rnd = 0; rnd < kRounds; rnd++) { + CUDA_CHECK(cudaMemcpy(d_c1, h_c1, ct_bytes, cudaMemcpyHostToDevice)); + rms[rnd] = paillier_randomize2( + d_c1, d_table, d_r, TAU_R, NUM_RND, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample); + } + printf("[Rand2] 计时完成\n\n"); + + // ── 7. D2H 最后一轮结果(d_c1 已是最后一轮输出)──────────────────── + { + struct timespec t0d, t1d; + clock_gettime(CLOCK_MONOTONIC, &t0d); + CUDA_CHECK(cudaMemcpy(h_c_out, d_c1, ct_bytes, cudaMemcpyDeviceToHost)); + clock_gettime(CLOCK_MONOTONIC, &t1d); + double d2h_ms = (double)(t1d.tv_sec - t0d.tv_sec) * 1e3 + + (double)(t1d.tv_nsec - t0d.tv_nsec) * 1e-6; + printf("[Rand2] D2H: %.3f ms 结果 %.2f MB 带宽 %.2f GB/s\n", d2h_ms, + ct_bytes / 1048576.0, + d2h_ms > 0.0 ? ct_bytes / 1e6 / d2h_ms : 0.0); + } + + /* + // ── 8. 写结果到 randomize2_results.txt ────────────────────────────── + // 格式(供 verify_randomize2.py 验证): + // 行1: NUM_RND ARR_LEN BASE_BITS TAU_R + // 行2: n² limbs(ARR_LEN 个 uint64,base-2^BASE_BITS,小端序) + // 行3: r_param = table[0] = (2^4352-1) mod n²(FMLM 恒等元) + // 行4: hs limbs(ARR_LEN 个 uint64,标准域) + // 每个用例 3 行: + // c̃₁ limbs(ARR_LEN 个 uint64,FMLM域) + // r_i limbs(TAU_R/64 个 uint64,64-bit packed,LSB在前) + // result limbs(ARR_LEN 个 uint64,FMLM域) + { + FILE *fout = fopen("randomize2_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Rand2] 无法创建 randomize2_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_RND, ARR_LEN, BASE_BITS, TAU_R); + // n² + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 rmax) rmax = rms[rnd]; + } + float ravg = rsum / kRounds; + printf("\n============================================================\n"); + printf(" paillier_randomize2 性能报告\n"); + printf(" 批大小=%d TAU_R=%d bit WINDOW=%d bit\n", NUM_RND, TAU_R, + WINDOW_BITS); + printf("============================================================\n"); + printf("[计时结果(cudaEvent,仅 GPU 核函数,不含传输)]\n"); + printf("------------------------------------------------------------\n"); + for (int rnd = 0; rnd < kRounds; rnd++) + printf(" 轮 %d : %8.3f ms (%.3f us/次)\n", rnd, rms[rnd], + rms[rnd] * 1e3f / NUM_RND); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %8.3f ms\n", ravg); + printf(" 最快轮 : %8.3f ms\n", rmin); + printf(" 最慢轮 : %8.3f ms\n", rmax); + printf(" 平均单次 Randomize2 : %8.3f us\n", ravg * 1e3f / NUM_RND); + printf(" (最快轮) 单次 : %8.3f us\n", rmin * 1e3f / NUM_RND); + printf("============================================================\n"); + + // ── 10. 释放 Randomize2 相关内存 ───────────────────────────────────── + cudaFreeHost(h_c1); + cudaFreeHost(h_r); + cudaFreeHost(h_c_out); + cudaFree(d_table); + cudaFree(d_c1); + cudaFree(d_r); + } + free(h_table); + printf("[Rand2] 测试完成\n"); +rand2_done:; + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_randomize2 测试完成。\n"); + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/paillier_subhomo.cu b/heu/library/algorithms/paillier_new/paillier_subhomo.cu new file mode 100644 index 0000000..73c0aa4 --- /dev/null +++ b/heu/library/algorithms/paillier_new/paillier_subhomo.cu @@ -0,0 +1,18804 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +#include +/// #include +#include + +#include + +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include +#include + +#include "cgbn/cgbn.h" + +// ============================================================================= +// CGBN 配置:用于 paillier_subhomo2 中的 GPU 批量模逆 +// 使用 inv_ 前缀避免与文件其他部分命名冲突 +// ============================================================================= +#define INV_BITS 4096 +#define INV_TPI 32 +typedef cgbn_context_t inv_context_t; +typedef cgbn_env_t inv_env_t; +typedef inv_env_t::cgbn_t inv_bn_t; +typedef cgbn_mem_t inv_bn_mem_t; + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 2048 // 指数 bit 数(对应 2048-bit Paillier λ) +#define EXP_U64_LIMBS (TAU / 64) // = 32,压缩格式每个指数占 32 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// FMLE_mod3_Kernel +// +// 功能:FMLE(x̃, e) → x^e mod n²(蒙哥马利域输入,普通域输出) +// +// 与 FMLMKernel_e 的区别: +// · 输入 A 已是蒙哥马利域(x̃ = x·R),故跳过 Step8(y←FMLM(y,r₁)) +// · 去掉参数 d_r1_ct / d_r1_nct(Step8 不再需要) +// · Step13 保留(t←FMLM(t,1)),将结果还原到普通域后输出 +// +// 与 FMLE_mod2_Kernel 的区别: +// · 保留 Step13,输出是普通域(x^e mod n²),而非蒙哥马利域 +// +// 参数: +// A [m * ARR_LEN] 输入,蒙哥马利域(x̃ = x·R) +// exp_bits [m * (tau/64)] 压缩格式指数,LSB-first,每指数 tau/64 个 +// uint64_t tau 指数 bit 长度 output [m * +// ARR_LEN] 输出,普通域(x^e mod n²) m 用例数量 d_r0 [ARR_LEN] +// r₀ = R mod n²(1 的蒙哥马利表示) 其余为 NTT 参数(与 FMLMKernel_e +// 完全一致) +// ============================================================= +__global__ void FMLE_mod3_Kernel( + const uint64_t *__restrict__ A, // 输入:x̃(蒙哥马利域) + const uint64_t + *__restrict__ exp_bits, // 指数(压缩格式,tau/64 个 uint64_t) + int tau, // 指数 bit 长度 + uint64_t *output, // 输出:x^e mod n²(普通域) + int m, // 用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1 的蒙哥马利表示) + // ★ 无 d_r1_ct / d_r1_nct:Step8 已跳过 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // 压缩格式加载指数:每 warp 的 lane 0~(exp_limbs-1) 各持有一个 uint64_t limb + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← x̃(直接加载蒙哥马利域输入) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原算法:y ← FMLM(y, r₁) = y·r²·R⁻¹ = y·R,将普通域转为蒙哥马利域 + // FMLE_mod3:输入 A 已是蒙哥马利域(x̃ = x·R),无需再乘 R,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9–12: 平方-乘循环(与 FMLMKernel_e 完全相同) + // + // 取第 i 个 bit:通过 __shfl_sync 从持有 limb(i/64) 的 lane 广播, + // 零额外全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + // Step 10: if e[i]=1 then t ← FMLM(t, y) + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // + // 去除蒙哥马利因子 R,将结果还原到普通域: + // t = (x^e · R) · 1 · R⁻¹ = x^e mod n² + // ★ 与 FMLE_mod2_Kernel 的关键区别:此步保留 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 200000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // FMLE 阶段耗时(各流并行,取最慢流) + float step4_ms; // XY 阶段耗时(各流并行,取最慢流) + float pipeline_ms; // Step③+④ 流水线实际墙钟时间(< step3+step4) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③+④ 流水线:每条流独立处理一个 chunk + // + // 旧方案(串行): + // Stream0..3: [FMLE₀][FMLE₁][FMLE₂][FMLE₃] → sync → [XY_all] + // + // 新方案(流水线): + // Stream 0: [FMLE₀ ───────────────────][XY₀ ───] + // Stream 1: [FMLE₁ ──────────────────] [XY₁ ──] + // Stream 2: [FMLE₂ ─────────────────] [XY₂ ─] + // Stream 3: [FMLE₃ ────────────────] [XY₃ ] + // + // XY₀ 可与 FMLE₁/₂/₃ 并发,节省 (NUM_STREAMS-1)×XY_chunk 时间 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + // ev_fmle_done[s]: 流 s 的 FMLE 完成时刻 + // ev_xy_done[s] : 流 s 的 XY 完成时刻 + cudaEvent_t ev_fmle_done[NUM_STREAMS]; + cudaEvent_t ev_xy_done[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_fmle_done[s])); + CUDA_CHECK(cudaEventCreate(&ev_xy_done[s])); + } + cudaEvent_t ev_pipe_start; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + // 在默认流上打起始时间戳,后续所有流从此刻开始计时 + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int fmle_blks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + // ★ 注意:XY 使用 actual 个 1-warp 块(32 线程/块)而非 ceil(actual/8) 个 + // 8-warp 块。 原因:若 actual 不能被 WARPS_PER_BLOCK=8 整除,ceil + // 会多出若干 warp, + // 这些 warp 会越界读写相邻流的 d_c1/d_hs_r 区域,造成竞争条件。 + // 使用 1-warp/块启动恰好产生 actual 个 warp,不产生越界访问。 + // XYfixWarpVector 使用 lane_id = threadIdx.x & 0x1F 做 warp 内索引, + // 与 blockDim.x=32 或 256 等价,结果完全正确。 + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + // ── Step③:FMLE(h̃_s^{r_i},蒙→蒙) ── + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + // 同一流上记录 FMLE 完成时间(无需额外同步,流内串行保证正确性) + CUDA_CHECK(cudaEventRecord(ev_fmle_done[s], streams[s])); + + // ── Step④:XY(c̃₁·h̃_s^{r_i},在同一流上立即接续) ── + // 同一流内核顺序执行,FMLE 完成后才会执行 XY,无需显式 Event 等待 + // actual 个块 × 32 线程/块 = 恰好 actual 个 warp,无越界 + XYfixWarpVector<<>>( + d_c1 + ct_off, d_hs_r + ct_off, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, + d_sample); + CUDA_CHECK(cudaEventRecord(ev_xy_done[s], streams[s])); + } + + // 等待所有流完成 + CUDA_CHECK(cudaDeviceSynchronize()); + CUDA_CHECK(cudaGetLastError()); + + // ── 统计各流耗时,取最慢流 ── + float max_fmle_ms = 0.f; + float max_xy_ms = 0.f; + float max_pipe_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float fmle_ms = 0.f, xy_ms = 0.f, total_ms = 0.f; + // FMLE 耗时:从流水线起点到该流 FMLE 完成 + CUDA_CHECK( + cudaEventElapsedTime(&fmle_ms, ev_pipe_start, ev_fmle_done[s])); + // XY 耗时:从该流 FMLE 完成到 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&xy_ms, ev_fmle_done[s], ev_xy_done[s])); + // 流水线总耗时:从起点到该流 XY 完成 + CUDA_CHECK(cudaEventElapsedTime(&total_ms, ev_pipe_start, ev_xy_done[s])); + if (fmle_ms > max_fmle_ms) max_fmle_ms = fmle_ms; + if (xy_ms > max_xy_ms) max_xy_ms = xy_ms; + if (total_ms > max_pipe_ms) max_pipe_ms = total_ms; + } + t.step3_ms = max_fmle_ms; + t.step4_ms = max_xy_ms; + t.pipeline_ms = max_pipe_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_fmle_done[s]); + cudaEventDestroy(ev_xy_done[s]); + } + cudaEventDestroy(ev_pipe_start); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// paillier_dec +// +// 功能:GPU 批量 Paillier 解密模幂(非 CRT 路线) +// output[i] = c[i]^λ mod n²(普通域) +// +// 输入 d_c_tilde 为蒙哥马利域密文 c̃ = c·R mod n², +// 调用 FMLE_mod3_Kernel:跳过 Step8(输入已是蒙哥马利域), +// 保留 Step13(将结果还原到普通域)。 +// +// 参数: +// d_c_tilde [batch × ARR_LEN] 输入:c̃(蒙哥马利域) +// d_exp [batch × tau/64] 指数 λ(压缩格式:每案例 tau/64 个 +// uint64_t) d_output [batch × ARR_LEN] 输出:c^λ mod n²(普通域) +// batch 批大小 +// tau λ 的 bit 长度 +// d_r0 [ARR_LEN] R mod n²(1 的蒙哥马利表示) +// 其余为 NTT 参数(与 paillier_mulplain 完全一致) +// +// 返回:GPU 端到端墙钟耗时(ms) +// ============================================================= +float paillier_dec( + const uint64_t *d_c_tilde, const uint64_t *d_exp, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + const int exp_u64_limbs = tau / 64; + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamCreate(&streams[s])); + + cudaEvent_t ev_start, ev_end; + CUDA_CHECK(cudaEventCreate(&ev_start)); + CUDA_CHECK(cudaEventCreate(&ev_end)); + CUDA_CHECK(cudaEventRecord(ev_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + const size_t ct_off = (size_t)offset * ARR_LEN; + const size_t exp_off = (size_t)offset * exp_u64_limbs; + + FMLE_mod3_Kernel<<>>( + d_c_tilde + ct_off, d_exp + exp_off, tau, d_output + ct_off, + actual_chunk, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + } + + CUDA_CHECK(cudaEventRecord(ev_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_end)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_start, ev_end)); + + for (int s = 0; s < NUM_STREAMS; s++) + CUDA_CHECK(cudaStreamDestroy(streams[s])); + CUDA_CHECK(cudaEventDestroy(ev_start)); + CUDA_CHECK(cudaEventDestroy(ev_end)); + + return ms; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float serial_ms = t.step3_ms + t.step4_ms; + float gpu_total = t.step1_ms + t.step2_ms + t.pipeline_ms; + float speedup = t.pipeline_ms > 0.f ? serial_ms / t.pipeline_ms : 0.f; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 各流模幂最大耗时 : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 各流乘法最大耗时 : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" Step3+4 无流水线参考值 : %8.4f ms\n", serial_ms); + printf(" Step3+4 流水线实际耗时 : %8.4f ms (加速 %.2fx)\n", + t.pipeline_ms, speedup); + printf(" ──────────────────────────────────────────\n"); + printf(" GPU 计算总时间(Step1+2+pipeline): %8.4f ms\n", gpu_total); + printf(" 平均每次 Randomize GPU 计算 : %8.4f us/次\n", + gpu_total * 1000.f / batch); + printf("======================================================\n"); +} + +// ============================================================================= +// paillier_subhomo2 所需的 GPU 模逆支持代码 +// ============================================================================= + +// ── ① CGBN 批量模逆 kernel(来自 invmod6.cu,使用 inv_ 前缀类型)───────────── +// +// 每个 CGBN Instance 由 INV_TPI=32 个线程协作处理一个 4096-bit 大数的模逆。 +// 输入:in_numbers[n] — 待求逆的大数数组(inv_bn_mem_t 格式,base-2^32) +// in_modulus — 共享模数 n²(一份,所有 instance 共用) +// 输出:out_inverses[n] — 模逆结果数组;若 gcd(x,m)>1 则置 0 +// ============================================================================= +__global__ void batch_mod_inverse_kernel_sub(cgbn_error_report_t *report, + inv_bn_mem_t *out_inverses, + inv_bn_mem_t *in_numbers, + inv_bn_mem_t *in_modulus, int n) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + int instance_id = tid / INV_TPI; + if (instance_id >= n) return; + + inv_context_t bn_context(cgbn_report_monitor, report, instance_id); + inv_env_t env(bn_context.env()); + + inv_bn_t x, m, r; + cgbn_load(env, x, in_numbers + instance_id); + cgbn_load(env, m, in_modulus); // 所有 instance 共用同一个 n² + + bool ok = cgbn_modular_inverse(env, r, x, m); + if (!ok) cgbn_set_ui32(env, r, 0); // gcd(x,m)>1:逆元不存在,填 0 + + cgbn_store(env, out_inverses + instance_id, r); +} + +// ── ② GPU 格式转换:base-2^17 uint64[] → inv_bn_mem_t(base-2^32 +// uint32)────── +// +// 逐实例:256 个 17-bit limbs(小端)→ 128 个 32-bit limbs(小端) +// 对应位段:output limb j 覆盖全局比特 [32j, 32j+31],最多跨 3 个 17-bit +// limbs。 启动参数:<<>>(每块 128 线程,每线程处理 1 个 32-bit +// 输出 limb) +// ============================================================================= +__global__ void format_to_cgbn_kernel( + inv_bn_mem_t *out, // [batch],CGBN 格式输出 + const uint64_t *in, // [batch × ARR_LEN],base-2^17 输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int j = threadIdx.x; // 32-bit 输出 limb 下标(0..127) + if (bid >= batch || j >= 128) return; + + const uint64_t *src = in + (size_t)bid * ARR_LEN; + uint32_t result = 0; + + // 收集从第 start_bit 开始的 32 个比特 + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; // 对应的 17-bit 起始 limb + int s_bit_off = start_bit % 17; // 在该 limb 内的偏移 + + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; // 当前 limb 剩余可用位数 + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[bid]._limbs[j] = result; +} + +// ── ③ GPU 格式转换:inv_bn_mem_t(base-2^32 uint32)→ base-2^17 uint64[]────── +// +// 逐实例:128 个 32-bit limbs(小端)→ 256 个 17-bit limbs(小端) +// 对应位段:output limb i 覆盖全局比特 [17i, 17i+16],最多跨 2 个 32-bit +// limbs。 高于 4095 比特的部分(bits 4096..4351,即 17-bit limbs 241..255 +// 的高位) 在合法 Paillier 密文中始终为 0,因此 src_limb≥128 时直接补 0 即可。 +// 启动参数:<<>>(每块 256 线程,每线程处理 1 个 17-bit 输出 limb) +// ============================================================================= +__global__ void format_from_cgbn_kernel( + uint64_t *out, // [batch × ARR_LEN],base-2^17 输出 + const inv_bn_mem_t *in, // [batch],CGBN 格式输入 + int batch) { + int bid = blockIdx.x; // 批次下标 + int i = threadIdx.x; // 17-bit 输出 limb 下标(0..255) + if (bid >= batch || i >= ARR_LEN) return; + + const uint32_t *src = in[bid]._limbs; + uint64_t result = 0; + + int start_bit = i * 17; + int s_limb = start_bit / 32; // 对应的 32-bit 起始 limb + int s_bit_off = start_bit % 32; // 在该 limb 内的偏移 + + int out_pos = 0; + int remaining = 17; + + while (remaining > 0) { + if (s_limb >= 128) break; // 超出 4096-bit 范围,高位恒为 0 + int available = 32 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = (src[s_limb] >> s_bit_off) & ((1u << take) - 1); + result |= ((uint64_t)piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + + out[(size_t)bid * ARR_LEN + i] = result; +} + +// ── ④ CPU 辅助:N2_arr(base-2^17 uint64,主机端)→ inv_bn_mem_t(base-2^32) +// +// 用于在 paillier_subhomo2 中将模数 n² 转换为 CGBN 格式后上传到设备。 +// ============================================================================= +static void bn17_to_cgbn_mem(const uint64_t *src, inv_bn_mem_t *dst) { + memset(dst->_limbs, 0, sizeof(dst->_limbs)); + for (int j = 0; j < 128; j++) { + uint32_t result = 0; + int start_bit = j * 32; + int out_pos = 0; + int remaining = 32; + int s_limb = start_bit / 17; + int s_bit_off = start_bit % 17; + while (remaining > 0 && s_limb < ARR_LEN) { + int available = 17 - s_bit_off; + int take = (available < remaining) ? available : remaining; + uint32_t piece = + (uint32_t)((src[s_limb] >> s_bit_off) & ((1u << take) - 1)); + result |= (piece << out_pos); + out_pos += take; + remaining -= take; + s_limb++; + s_bit_off = 0; + } + dst->_limbs[j] = result; + } +} + +// ============================================================================= +// SubHomo 辅助:base-2^17 uint64_t[] ↔ libtommath mp_int +// ============================================================================= + +// uint64_t[ARR_LEN] base-2^17 小端序 → mp_int +static void bn17_to_mp_u64(const uint64_t *d17, mp_int *out) { + // ARR_LEN*BASE_BITS = 256*17 = 4352 bit = 544 字节 + uint8_t le_bytes[544]; + memset(le_bytes, 0, sizeof(le_bytes)); + for (int i = 0; i < ARR_LEN; i++) { + uint32_t val = (uint32_t)(d17[i] & 0x1FFFFu); + int bit_base = i * BASE_BITS; + for (int b = 0; b < BASE_BITS; b++) { + if (val & (1u << b)) { + int pos = bit_base + b; + le_bytes[pos >> 3] |= (uint8_t)(1u << (pos & 7)); + } + } + } + int used = 544; + while (used > 1 && le_bytes[used - 1] == 0) used--; + uint8_t be_bytes[544]; + for (int i = 0; i < used; i++) be_bytes[i] = le_bytes[used - 1 - i]; + mp_from_ubin(out, be_bytes, (size_t)used); +} + +// mp_int → uint64_t[ARR_LEN] base-2^17 小端序 +static void mp_to_bn17_u64(const mp_int *src, uint64_t *d17) { + memset(d17, 0, ARR_LEN * sizeof(uint64_t)); + size_t needed = (size_t)mp_ubin_size(src); + if (needed == 0) return; + if (needed > 544) needed = 544; + uint8_t be_bytes[544]; + memset(be_bytes, 0, sizeof(be_bytes)); + // mp_to_ubin 输出大端序,写入 be_bytes 末尾 needed 字节 + mp_to_ubin(src, be_bytes + (544 - needed), needed, NULL); + // 翻转为小端序字节数组 + uint8_t le_bytes[544]; + for (int i = 0; i < 544; i++) le_bytes[i] = be_bytes[543 - i]; + // 提取 17-bit limbs + for (int i = 0; i < ARR_LEN; i++) { + int bit_base = i * BASE_BITS; + uint32_t val = 0; + for (int b = 0; b < BASE_BITS; b++) { + int pos = bit_base + b; + if (le_bytes[pos >> 3] & (1u << (pos & 7))) val |= (1u << b); + } + d17[i] = (uint64_t)val; + } +} + +// ============================================================================= +// paillier_subhomo:密文同态减法 c̃ = c̃₁ · c₂⁻¹·R (mod n²) +// +// SubHomo(c̃₁, c̃₂) 四步流程: +// ① GPU XYfixWarpIRVector <<>>: c̃₂ → c₂ 蒙哥马利 → 普通域 +// ← D2H +// ② CPU mp_invmod(c₂, n²) → c₂⁻¹ 逐元素 CPU 求逆 +// → H2D +// ③ GPU XYfixWarpROneVector<<>>: c₂⁻¹ → c̃₂⁻¹ 普通域 → 蒙哥马利 +// ④ GPU XYfixWarpVector <<>>: c̃₁·c̃₂⁻¹ → c̃ 蒙哥马利乘法 +// +// 注:XYfixWarpROneVector 按 i=blockIdx.x*32+threadIdx.x 同时索引 inoutAct, +// 故 d_r1_ct / d_r1_nct 必须先 broadcast 成 batch×ARR_LEN 份再传入。 +// ============================================================================= +void paillier_subhomo( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr // [ARR_LEN] n²(主机端,base-2^17 小端序) +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 原地改为普通域 c₂ + uint64_t *d_inv_dev = NULL; // Step②→③→④ 的 c₂⁻¹ / c̃₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + // 主机 pinned 缓冲区 + uint64_t *h_c2 = NULL; // D2H:普通域 c₂ + uint64_t *h_inv = NULL; // CPU 计算的 c₂⁻¹,待 H2D + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_c2, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_inv, batch_bytes)); + + // 将 c̃₁ 复制到 d_result(Step④ 将原地用乘积覆盖) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 将原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + printf("[SubHomo] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── D2H:普通域 c₂ 传回主机 ──────────────────────────────────────────────── + printf("[SubHomo] D2H c₂...\n"); + CUDA_CHECK(cudaMemcpy(h_c2, d_c2_work, batch_bytes, cudaMemcpyDeviceToHost)); + + // ── Step②:CPU c₂⁻¹ = mp_invmod(c₂, n²) ────────────────────────────────── + printf("[SubHomo] Step② CPU mp_invmod ×%d...\n", batch); + { + mp_int n2; + mp_init(&n2); + bn17_to_mp_u64(N2_arr, &n2); + + int ok = 0, skip = 0; + for (int k = 0; k < batch; k++) { + const uint64_t *c2_limbs = h_c2 + (size_t)k * ARR_LEN; + uint64_t *iv_limbs = h_inv + (size_t)k * ARR_LEN; + mp_int c2, inv_c2; + mp_init(&c2); + mp_init(&inv_c2); + bn17_to_mp_u64(c2_limbs, &c2); + if (mp_invmod(&c2, &n2, &inv_c2) == MP_OKAY) { + mp_to_bn17_u64(&inv_c2, iv_limbs); + ok++; + } else { + // gcd(c₂, n²) > 1:逆元不存在,填零(该 case result 无意义) + memset(iv_limbs, 0, ARR_LEN * sizeof(uint64_t)); + skip++; + } + mp_clear(&c2); + mp_clear(&inv_c2); + } + mp_clear(&n2); + printf("[SubHomo] Step② 完成:ok=%d, skip(gcd>1)=%d\n", ok, skip); + } + + // ── H2D:c₂⁻¹ 上传至设备 + // ──────────────────────────────────────────────────── + printf("[SubHomo] H2D c₂⁻¹...\n"); + CUDA_CHECK(cudaMemcpy(d_inv_dev, h_inv, batch_bytes, cudaMemcpyHostToDevice)); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + printf("[SubHomo] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + printf("[SubHomo] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFreeHost(h_c2); + cudaFreeHost(h_inv); +} + +// ============================================================================= +// paillier_subhomo2:密文同态减法(全 GPU,无 D2H/H2D 传输) +// +// 与 paillier_subhomo 接口完全相同,但将: +// D2H → CPU mp_invmod(串行)→ H2D +// 替换为: +// GPU format_to_cgbn → GPU batch_mod_inverse_kernel_sub(CGBN 并行) +// → GPU format_from_cgbn +// 从而消除主机-设备数据传输,所有步骤在 GPU 上连续完成。 +// +// SubHomo2(c̃₁, c̃₂) 完整流程: +// ① GPU XYfixWarpIRVector <<>> : c̃₂ → c₂ (蒙哥马利→普通域) +// ②a GPU format_to_cgbn_kernel<<>> : c₂ → c₂_cgbn +// (base-2^17→CGBN) ②b GPU batch_mod_inverse_kernel_sub : c₂_cgbn → +// inv_cgbn (CGBN 并行模逆) ②c GPU format_from_cgbn_kernel<<>>: +// inv_cgbn → c₂⁻¹ (CGBN→base-2^17) ③ GPU XYfixWarpROneVector<<>>: +// c₂⁻¹ → c̃₂⁻¹ (普通域→蒙哥马利) ④ GPU XYfixWarpVector <<>>: +// c̃₁·c̃₂⁻¹ → c̃ (蒙哥马利乘法) +// ============================================================================= +void paillier_subhomo2( + int batch, + const uint64_t *d_c1_tilde, // [batch×ARR_LEN] 输入 c̃₁(只读) + const uint64_t *d_c2_tilde, // [batch×ARR_LEN] 输入 c̃₂(只读) + uint64_t *d_result, // [batch×ARR_LEN] 输出 c̃₁·c̃₂⁻¹ + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample, + const uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式 + const uint64_t *N2_arr, // [ARR_LEN] n²(主机端,base-2^17 小端序) + bool verbose // 是否打印进度 +) { + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + // ── 设备工作缓冲区 ───────────────────────────────────────────────────────── + uint64_t *d_c2_work = NULL; // c̃₂ 可写副本;Step① 后变为普通域 c₂ + inv_bn_mem_t *d_c2_cgbn = NULL; // c₂ 的 CGBN (base-2^32) 格式 + inv_bn_mem_t *d_inv_cgbn = NULL; // 模逆结果的 CGBN 格式 + uint64_t *d_inv_dev = NULL; // 模逆结果转回 base-2^17 格式:c₂⁻¹ + uint64_t *d_r1_ct_bat = NULL; // R² CT 广播后的 batch 份 + uint64_t *d_r1_nct_bat = NULL; // R² NCT 广播后的 batch 份 + inv_bn_mem_t *d_n2_cgbn = NULL; // n² 的 CGBN 格式(模数,单份共享) + + CUDA_CHECK(cudaMalloc(&d_c2_work, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_cgbn, (size_t)batch * sizeof(inv_bn_mem_t))); + CUDA_CHECK(cudaMalloc(&d_inv_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_ct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_r1_nct_bat, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_n2_cgbn, sizeof(inv_bn_mem_t))); + + // ── 将 n²(host, base-2^17)转换为 CGBN 格式并上传设备 ─────────────────── + { + inv_bn_mem_t h_n2_cgbn; + bn17_to_cgbn_mem(N2_arr, &h_n2_cgbn); + CUDA_CHECK(cudaMemcpy(d_n2_cgbn, &h_n2_cgbn, sizeof(inv_bn_mem_t), + cudaMemcpyHostToDevice)); + } + + // 将 c̃₁ 复制到 d_result(Step④ 原地覆盖为乘积) + CUDA_CHECK( + cudaMemcpy(d_result, d_c1_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + // 将 c̃₂ 复制到工作区(Step① 原地改为普通域) + CUDA_CHECK( + cudaMemcpy(d_c2_work, d_c2_tilde, batch_bytes, cudaMemcpyDeviceToDevice)); + + // 将单份 R² CT/NCT 广播为 batch×ARR_LEN(供 Step③ 使用) + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_r1_ct_bat, d_r1_ct, batch); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_r1_nct_bat, d_r1_nct, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── Step①:GPU c̃₂ → c₂ = FMLM(c̃₂, 1) 蒙哥马利 → 普通域 ─────────────── + if (verbose) printf("[SubHomo2] Step① GPU c̃₂→c₂,batch=%d\n", batch); + XYfixWarpIRVector<<>>( + d_c2_work, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②a:GPU c₂(base-2^17 uint64)→ c₂_cgbn(CGBN base-2^32)───────── + // 每个 batch item 对应一个 block(128 线程,各处理一个 32-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②a 格式转换 base-2^17 → CGBN\n"); + format_to_cgbn_kernel<<>>(d_c2_cgbn, d_c2_work, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step②b:GPU CGBN 批量模逆 c₂_cgbn → inv_cgbn ───────────────────────── + // TPI=32 线程/instance;threads_per_block=256 → 8 instances/block + if (verbose) printf("[SubHomo2] Step②b GPU CGBN 批量模逆,batch=%d\n", batch); + { + cgbn_error_report_t *report; + cgbn_error_report_alloc(&report); + + const int threads_per_block = 256; // 256 / INV_TPI = 8 instances/block + const int blocks = + (batch * INV_TPI + threads_per_block - 1) / threads_per_block; + batch_mod_inverse_kernel_sub<<>>( + report, d_inv_cgbn, d_c2_cgbn, d_n2_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + if (cgbn_error_report_check(report)) + fprintf(stderr, "[SubHomo2] CGBN 内部错误: %s\n", + cgbn_error_string(report)); + cgbn_error_report_free(report); + } + + // ── Step②c:GPU inv_cgbn(CGBN)→ c₂⁻¹(base-2^17 uint64)────────────── + // 每个 batch item 一个 block(256 线程,各处理一个 17-bit 输出 limb) + if (verbose) printf("[SubHomo2] Step②c 格式转换 CGBN → base-2^17\n"); + format_from_cgbn_kernel<<>>(d_inv_dev, d_inv_cgbn, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step③:GPU c₂⁻¹ → c̃₂⁻¹ = FMLM(c₂⁻¹, R²) 普通域 → 蒙哥马利 ───────── + if (verbose) printf("[SubHomo2] Step③ GPU c₂⁻¹→c̃₂⁻¹,batch=%d\n", batch); + XYfixWarpROneVector<<>>( + d_inv_dev, d_r1_ct_bat, d_r1_nct_bat, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── Step④:GPU c̃ = FMLM(c̃₁, c̃₂⁻¹) 蒙哥马利乘法 ────────────────────────── + if (verbose) printf("[SubHomo2] Step④ GPU c̃₁·c̃₂⁻¹→c̃,batch=%d\n", batch); + XYfixWarpVector<<>>( + d_result, // inout = c̃₁(Step 前已复制,结果原地覆盖) + d_inv_dev, // inoutA = c̃₂⁻¹ + d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + + // ── 释放工作缓冲区 ───────────────────────────────────────────────────────── + cudaFree(d_c2_work); + cudaFree(d_c2_cgbn); + cudaFree(d_inv_cgbn); + cudaFree(d_inv_dev); + cudaFree(d_r1_ct_bat); + cudaFree(d_r1_nct_bat); + cudaFree(d_n2_cgbn); +} + +// ───────────────────────────────────────────────────────────────────────────── +#define WINDOW_BITS 10 +#define TABLE_SIZE (1 << WINDOW_BITS) // 1024 个表项,对应 10-bit 窗口 + +// ───────────────────────────────────────────────────────────────────────────── + +// ============================================================================= +// generate_hs_table +// +// 随机生成一个约 4096-bit 的底数 hs(满足 0 < hs < n²), +// 建立 WINDOW_BITS=10 bit 窗口的模幂预计算表(FMLM 域,CPU 端): +// h_table[i × ARR_LEN] = hs^i · r mod n², i = 0, 1, …, TABLE_SIZE-1 +// +// 其中 r = 2^(ARR_LEN × BASE_BITS) - 1 = 2^4352 - 1,是 FMLM 算法(Algorithm +// 1) 的参数,与 XYfixWarpVector 实现的乘法一致。 FMLM 域恒等元为 r mod +// n²(即标准域中 1 的 FMLM 表示)。 +// +// 递推关系(在标准大整数乘法意义下,非 FMLM 乘法): +// table[0] = r mod n² +// table[i] = table[i-1] × hs mod n² → = hs^i · r mod n² ✓ +// +// 参数: +// N2_arr [ARR_LEN] n²(base-2^17 小端序,主机端只读) +// h_hs [ARR_LEN] 输出:随机底数 hs(base-2^17,标准域) +// h_table [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域,base-2^17) +// ============================================================================= +void generate_hs_table( + const uint64_t *N2_arr, // [ARR_LEN] n²(只读) + uint64_t *h_hs, // [ARR_LEN] 输出:底数 hs + uint64_t *h_table // [TABLE_SIZE × ARR_LEN] 输出:预计算表(FMLM 域) +) { + // ── 1. 随机生成 hs ∈ (0, n²) + // ──────────────────────────────────────────────── n² 最高非零 limb 在索引 + // 240,值为 27958(约 4095 bit) 策略: + // limb[0..239] ← 随机 17-bit 值(均小于 BASE=131072,满足 < n²) + // limb[240] ← rand() % n2_top_val(保证 hs < n²),重试到非零 + // limb[241..] ← 0 + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + do { + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) { + h_hs[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & + mask17; + } else if (j == n2_top_idx) { + h_hs[j] = (uint64_t)(rand() % (int)n2_top_val); + } else { + h_hs[j] = 0ULL; + } + } + } while (h_hs[n2_top_idx] == 0); // 极小概率(< 1/27958)需要重试 + + printf( + "[generate_hs_table] 随机底数 hs 生成完成" + "(limb[%d] = %llu,约 %d bit)\n", + n2_top_idx, (unsigned long long)h_hs[n2_top_idx], + n2_top_idx * BASE_BITS + 14); // 14 ≈ floor(log2(27958)) + + // ── 2. 转换 n² 和 hs 为 mp_int + // ────────────────────────────────────────────── + mp_int n2_mp, base_mp, cur_mp, r_mp; + mp_init(&n2_mp); + mp_init(&base_mp); + mp_init(&cur_mp); + mp_init(&r_mp); + + bn17_to_mp_u64(N2_arr, &n2_mp); + bn17_to_mp_u64(h_hs, &base_mp); + + // ── 3. 计算 FMLM 参数 r = 2^(ARR_LEN×BASE_BITS) - 1 = 2^4352 - 1,再对 n² + // 取模 r mod n² 是 FMLM 域中 1 的表示(恒等元),将作为预计算表的第 0 项 + mp_2expt(&r_mp, (int)ARR_LEN * BASE_BITS); // r_mp = 2^4352 + mp_sub_d(&r_mp, 1, &r_mp); // r_mp = 2^4352 - 1 + mp_mod(&r_mp, &n2_mp, &r_mp); // r_mp = (2^4352 - 1) mod n² + + printf( + "[generate_hs_table] FMLM 参数 r = 2^4352 - 1," + "r mod n² bit 长度 = %d\n", + mp_count_bits(&r_mp)); + + // ── 4. 建立预计算表:table[i] = hs^i · r mod n²(FMLM 域)────────────────── + printf( + "[generate_hs_table] 开始建立预计算表" + "(WINDOW=%d bit,TABLE_SIZE=%d 项)...\n", + WINDOW_BITS, TABLE_SIZE); + + struct timespec ts0, ts1; + clock_gettime(CLOCK_MONOTONIC, &ts0); + + // table[0] = r mod n²(FMLM 域恒等元,即 hs^0 · r mod n²) + mp_copy(&r_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)0 * ARR_LEN); + + // table[i] = table[i-1] × hs mod n²(普通乘法,递推得 hs^i · r mod n²) + // mp_mulmod 支持输出与输入别名(内部创建临时数),故原地更新 cur_mp 安全 + for (int i = 1; i < TABLE_SIZE; i++) { + mp_mulmod(&cur_mp, &base_mp, &n2_mp, &cur_mp); + mp_to_bn17_u64(&cur_mp, h_table + (size_t)i * ARR_LEN); + } + + clock_gettime(CLOCK_MONOTONIC, &ts1); + double elapsed_ms = (double)(ts1.tv_sec - ts0.tv_sec) * 1.0e3 + + (double)(ts1.tv_nsec - ts0.tv_nsec) * 1.0e-6; + + printf( + "[generate_hs_table] 预计算表构建完成" + "(%d 次 mp_mulmod,耗时 %.1f ms,平均 %.3f ms/次)\n", + TABLE_SIZE - 1, elapsed_ms, elapsed_ms / (TABLE_SIZE - 1)); + + // ── 5. 清理 + // ────────────────────────────────────────────────────────────────── + mp_clear(&n2_mp); + mp_clear(&base_mp); + mp_clear(&cur_mp); + mp_clear(&r_mp); +} + +// ============================================================================= +// Getrn +// +// GPU 内核:对一批密文,各自计算 hs^{r_i} mod n²,输出蒙哥马利域(FMLM +// 域)结果。 +// +// 算法:左到右 WINDOW_BITS-bit 窗口模幂(Left-to-Right Windowed +// Exponentiation) +// +// 每个 warp 独立处理一个密文,执行: +// t = table[0] ← FMLM 域恒等元(r mod n²) +// for k = num_windows-1 downto 0: +// t = Square^{WINDOW_BITS}(t) ← WINDOW_BITS 次 FMLM 平方 +// w = r 的第 [k·W, (k+1)·W-1] 位组成的 WINDOW_BITS-bit 窗口值 +// t = FMLM(t, table[w]) ← 查表乘法 +// output = t ← hs^r · r^{-1} · r = hs^r·r mod n²(FMLM +// 域) +// +// 正确性说明: +// table[i] = hs^i · r mod n²(generate_hs_table 建立,FMLM 域) +// table[0] = r mod n²(FMLM 恒等元,FMLM(identity, X) = X) +// 对平方:FMLM(ã, ã) = a²·r mod n²,即 a² 的 FMLM 表示 ✓ +// 对乘法:FMLM(ã, table[w]) = a·hs^w·r mod n²,即 (a·hs^w) 的 FMLM 表示 ✓ +// 最终 t = hs^r · r mod n² ∈ FMLM 域 ✓ +// +// 参数: +// d_table [TABLE_SIZE × ARR_LEN] generate_hs_table 输出的 FMLM 域预计算表 +// d_r_batch [batch × exp_limbs] 随机指数 r,压缩格式(小端序 uint64 +// limbs) tau r 的 bit 长度(须为 64 +// 的整数倍) d_output [batch × ARR_LEN] 输出 hs^r(FMLM +// 域,base-2^17) batch 密文数量(每个 warp +// 处理一个) +// +// 共享内存布局(每个 warp 512 个 uint64): +// ws[0..255] = t_buf (FMLM 域累加器) +// ws[256..511] = entry_buf (从全局内存加载的当前表项) +// +// 启动参数(示例,WARP_PER_BLK=8): +// smem_size = WARP_PER_BLK * 512 * sizeof(uint64_t) +// Getrn<<<(batch+WARP_PER_BLK-1)/WARP_PER_BLK, WARP_PER_BLK*32, +// smem_size>>>(...) +// ============================================================================= +__global__ void Getrn( + const uint64_t + *__restrict__ d_table, // [TABLE_SIZE × ARR_LEN] FMLM 域预计算表 + const uint64_t + *__restrict__ d_r_batch, // [batch × exp_limbs] 随机指数(压缩格式) + int tau, // r 的 bit 长度 + uint64_t *d_output, // [batch × ARR_LEN] 输出 hs^r(FMLM 域) + int batch, // 密文数量 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= batch) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const t_buf = ws; // [256] FMLM 域累加器 + uint64_t *const entry_buf = ws + 256; // [256] 当前查表项 + + const int off = global_warp_id * ARR_LEN; + + // ── 加载本 warp 的指数 r(压缩格式:exp_limbs 个 uint64,小端序)── + const int exp_limbs = tau / 64; + const uint64_t *my_exp = d_r_batch + (size_t)global_warp_id * exp_limbs; + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_exp[lane_id]; + +// ── 初始化 t = table[0](FMLM 域恒等元 = r mod n²,即 hs^0 的 FMLM 表示)── +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_table[8 * lane_id + j]; + __syncwarp(); + + // ── 左到右 WINDOW_BITS-bit 窗口模幂 + // ────────────────────────────────────────── + // + // num_windows = ceil(tau / WINDOW_BITS) + // 窗口 k(k = num_windows-1 .. 0)覆盖指数 bit [k*W .. (k+1)*W-1](LSB 编号) + // 第一个窗口(最高位)对 identity 做 WINDOW_BITS 次平方仍得 identity, + // 随后乘以 table[top_window_val],自然进入正确递推,无需特判。 + // + const int num_windows = (tau + WINDOW_BITS - 1) / WINDOW_BITS; + + for (int k = num_windows - 1; k >= 0; k--) { +// ① WINDOW_BITS 次 FMLM 平方:t ← t^{2^WINDOW_BITS}(FMLM 域) +// 展开 WINDOW_BITS=10 次,避免循环开销;每次平方后同步 warp +#pragma unroll + for (int sq = 0; sq < WINDOW_BITS; sq++) { + XYfixWarpSquareVector_dev(t_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ② 提取 10-bit 窗口值 + // bit_pos = k * WINDOW_BITS:窗口起始位(LSB 编号) + // 每个 lane 通过 __shfl 读取覆盖该窗口的一或两个 uint64 limb, + // warp 内所有 lane 计算出相同的 window_val(无分歧) + const int bit_pos = k * WINDOW_BITS; + const int limb_idx = bit_pos / 64; + const int bit_off = bit_pos % 64; + + uint64_t lo = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)limb_idx); + uint64_t window_val = (lo >> bit_off) & ((1ULL << WINDOW_BITS) - 1ULL); + + // 窗口跨 limb 边界:从相邻高位 limb 补充剩余位 + if (bit_off + WINDOW_BITS > 64 && limb_idx + 1 < exp_limbs) { + uint64_t hi = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(limb_idx + 1)); + window_val |= (hi << (64 - bit_off)) & ((1ULL << WINDOW_BITS) - 1ULL); + } + + // ③ 加载 table[window_val] 到 entry_buf(全 warp 协作,coalesced 读取) + // window_val 对全 warp 统一,d_table 偏移确定,无 bank conflict + const uint64_t *entry = d_table + window_val * (uint64_t)ARR_LEN; +#pragma unroll + for (int j = 0; j < 8; j++) + entry_buf[8 * lane_id + j] = entry[8 * lane_id + j]; + __syncwarp(); + + // ④ FMLM 乘法:t ← FMLM(t, table[window_val])(结果写回 t_buf) + XYfixWarpVector_dev(t_buf, entry_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + +// ── 写回全局内存(t_buf 即 hs^r 的 FMLM 域表示)── +#pragma unroll + for (int j = 0; j < 8; j++) + d_output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// paillier_randomize2 +// +// 功能:批量随机化密文,执行 Randomize(c̃₁) = c̃₁ · h̃_s^{r_i} mod n²(FMLM域) +// +// 流程: +// Step① Getrn 内核:利用 10-bit 窗口预计算表计算 h̃_s^{r_i}(FMLM域输出) +// Step② XYfixWarpVector:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) +// +// 参数: +// d_c1 [batch × ARR_LEN] 输入 c̃₁(FMLM域),结果原地覆写 +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域,由 +// generate_hs_table 建立) d_r_batch [batch × tau_r/64] 随机指数 +// r_i(压缩格式,小端序 uint64) tau_r r 的 +// bit 长度(Paillier 加密随机数为 k/2 = 1024 bit) batch 密文数量 其余为 NTT +// 参数(与 Getrn / XYfixWarpVector 完全一致) +// +// 返回:Step①+② GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_randomize2( + uint64_t *d_c1, // [batch × ARR_LEN] 输入兼输出(FMLM域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条密文,WARP_PER_BLK 个 warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:c̃₁[i] · h̃_s^{r_i} → 结果原地写回 d_c1(FMLM域) + XYfixWarpVector<<>>( + d_c1, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +// ============================================================= +// paillier_encryptionhs +// +// 功能:批量 Paillier 加密 +// c̃_out[i] = FMLM( (1+m_i·n), h̃_s^{r_i} ) mod n²(FMLM 域) +// +// 流程: +// Step① Getrn:10-bit 窗口查表计算 h̃_s^{r_i}(FMLM域输出) +// Step② compute_plain_factor_batch_kernel:(1+m·n) mod n²(标准域) +// Step③ XYfixWarpROneVector:(1+m·n) → FMLM 域,即 (1+m·n)·R mod n² +// Step④ XYfixWarpVector:c̃ = FMLM((1̃+m·n), h̃_s^{r_i}) +// +// 正确性验证: +// c̃ = (1+m·n)·h_s^r·R mod n² (FMLM 表示) +// 还原:c_std = c̃ · R⁻¹ mod n² = (1+m·n)·h_s^r mod n² +// +// 参数: +// d_m_batch [batch × ARR_LEN] 明文 m(base-2^BASE_BITS,标准域) +// d_ct_out [batch × ARR_LEN] 输出密文(FMLM域) +// d_n [ARR_LEN] 公钥 n(标准域) +// d_n2 [ARR_LEN] 模数 n²(标准域) +// d_table [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) +// d_r_batch [batch × tau_r/64] 随机指数 r_i(小端序 uint64) +// tau_r r 的 bit 长度(Paillier 中 = 1024) +// batch 明文数量 +// d_r1_ct [ARR_LEN] R² mod n² 的 CT 形式(所有 warp +// 共享) d_r1_nct [ARR_LEN] R² mod n² 的 NCT 形式(所有 +// warp 共享) 其余为 NTT 参数(与 Getrn / XYfixWarpVector 一致) +// +// 返回:Step①~④ GPU 总耗时(ms,cudaEvent 精度) +// ============================================================= +float paillier_encryptionhs( + const uint64_t *d_m_batch, // [batch × ARR_LEN] 明文(标准域) + uint64_t *d_ct_out, // [batch × ARR_LEN] 输出密文(FMLM域) + const uint64_t *d_n, // [ARR_LEN] 公钥 n(标准域) + const uint64_t *d_n2, // [ARR_LEN] 模数 n²(标准域) + const uint64_t *d_table, // [TABLE_SIZE × ARR_LEN] hs 预计算表(FMLM域) + const uint64_t *d_r_batch, // [batch × tau_r/64] 随机指数 r_i + int tau_r, // r 的 bit 长度(= 1024) + int batch, + const uint64_t + *d_r1_ct, // [ARR_LEN] R² CT 形式(XYfixWarpROneVector 用) + const uint64_t + *d_r1_nct, // [ARR_LEN] R² NCT 形式(XYfixWarpROneVector 用) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + // 分配临时 buffer 存储 Getrn 输出 h̃_s^{r_i}(FMLM域) + uint64_t *d_hs_r_tilde = nullptr; + CUDA_CHECK( + cudaMalloc(&d_hs_r_tilde, (size_t)batch * ARR_LEN * sizeof(uint64_t))); + + // Getrn 启动参数:每 warp 处理一条加密,WARP_PER_BLK warp 一个 block + const int g_wpb = WARP_PER_BLK; + const int g_blks = (batch + g_wpb - 1) / g_wpb; + const size_t g_smem = (size_t)g_wpb * 512 * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + CUDA_CHECK(cudaEventRecord(e0)); + + // Step①:Getrn 计算 h̃_s^{r_i}(FMLM域),结果写入 d_hs_r_tilde + Getrn<<>>( + d_table, d_r_batch, tau_r, d_hs_r_tilde, batch, d_negmodn, + d_negmodn_shoup, d_modn, d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, + d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, + d_NCTtwiddle_shoup, d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, + inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step②:计算 (1 + m·n) mod n²(标准域),结果写入 d_ct_out + // 每个 block 处理一个明文,blockDim.x = ARR_LEN 线程各负责一个 limb + compute_plain_factor_batch_kernel<<>>( + d_m_batch, d_n, d_n2, + d_ct_out, // 复用输出缓冲区暂存 (1+m·n)(标准域) + +1 // sign = +1:(1 + m·n) + ); + CUDA_CHECK(cudaGetLastError()); + + // Step③:(1 + m·n) 转入 FMLM 域 + // XYfixWarpROneVector 计算 FMLM(factor, R²) = factor·R mod n² + // 注意:XYfixWarpROneVector 用 lane_id(而非 blockIdx.x)索引 + // inoutAct/inoutAnct, + // 故 d_r1_ct / d_r1_nct 只需单份 ARR_LEN 数组,所有 warp 共享 + XYfixWarpROneVector<<>>( + d_ct_out, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + // Step④:FMLM((1̃+m·n), h̃_s^{r_i}) → 密文 c̃(FMLM域),原地写回 d_ct_out + // FMLM(A, B) = A·B·R⁻¹ mod n² + // = (1+m·n)·R · h_s^r·R · R⁻¹ = (1+m·n)·h_s^r·R mod n² ✓ + XYfixWarpVector<<>>( + d_ct_out, d_hs_r_tilde, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, e0, e1)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + cudaFree(d_hs_r_tilde); + return ms; +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + /* + // ============================================================ + // [paillier_dec 性能测试] + // kWarmup 次预热 + kRounds 次计时,每轮批量处理 NUM_DEC 个密文, + // 计算 c^λ mod n²(非 CRT),输出结果到 dec_results.txt 供 Python 验证。 + // ============================================================ + const int NUM_DEC = 1000; + const int kWarmup = 3; + const int kRounds = 5; + const size_t ct_bytes = (size_t)NUM_DEC * ARR_LEN * sizeof(uint64_t); + const size_t exp_bytes = (size_t)NUM_DEC * EXP_U64_LIMBS * sizeof(uint64_t); + + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + // ── 1. 主机 pinned 内存 ──────────────────────────────────────────────────── + uint64_t *h_ct = NULL; + uint64_t *h_exp = NULL; + uint64_t *h_result = NULL; + CUDA_CHECK(cudaMallocHost(&h_ct, ct_bytes)); + CUDA_CHECK(cudaMallocHost(&h_exp, exp_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, ct_bytes)); + + // ── 2. 随机生成蒙哥马利密文 c̃(范围 [0, + n²))──────────────────────────────── srand((unsigned int)time(NULL)); for + (int p = 0; p < NUM_DEC; p++) { uint64_t *c = h_ct + (size_t)p * ARR_LEN; for + (int j = 0; j < ARR_LEN; j++) { if (j < n2_top_idx) { c[j] = + ((uint64_t)(uint32_t)rand() ^ ((uint64_t)(uint32_t)rand() << 15)) & mask17; } + else if (j == n2_top_idx) { c[j] = (uint64_t)(rand() % (int)n2_top_val); } + else { c[j] = 0ULL; + } + } + } + printf("[Dec] 随机生成 %d 个蒙哥马利域密文完成\n", NUM_DEC); + + // ── 3. 随机生成 1 个 TAU-bit λ,广播到所有用例 ─────────────────────────── + // λ 以压缩格式存储:EXP_U64_LIMBS 个 uint64_t,LSB 在前 + uint64_t lambda_limbs[EXP_U64_LIMBS]; + for (int j = 0; j < EXP_U64_LIMBS; j++) + lambda_limbs[j] = ((uint64_t)(uint32_t)rand() << 32) | + (uint64_t)(uint32_t)rand(); for (int p = 0; p < NUM_DEC; p++) memcpy(h_exp + + (size_t)p * EXP_U64_LIMBS, lambda_limbs, EXP_U64_LIMBS * sizeof(uint64_t)); + printf("[Dec] 随机生成 λ(TAU=%d bit),广播到全部 %d 个用例\n", TAU, + NUM_DEC); + + // ── 4. 设备内存分配 & H2D ───────────────────────────────────────────────── + uint64_t *d_ct = NULL; + uint64_t *d_exp_d = NULL; + uint64_t *d_result = NULL; + CUDA_CHECK(cudaMalloc(&d_ct, ct_bytes)); + CUDA_CHECK(cudaMalloc(&d_exp_d, exp_bytes)); + CUDA_CHECK(cudaMalloc(&d_result, ct_bytes)); + CUDA_CHECK(cudaMemcpy(d_ct, h_ct, ct_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaMemcpy(d_exp_d, h_exp, exp_bytes, cudaMemcpyHostToDevice)); + + // ── 5. 预热 ─────────────────────────────────────────────────────────────── + printf("[Dec] 预热 %d 次(不计时)...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + } + printf("[Dec] 预热完成\n\n"); + + // ── 6. 计时 ─────────────────────────────────────────────────────────────── + printf("[Dec] 开始计时(%d 轮 × %d 个密文)...\n", kRounds, NUM_DEC); + double round_ms[5]; + for (int r = 0; r < kRounds; r++) { + struct timespec t0, t1; + clock_gettime(CLOCK_MONOTONIC, &t0); + paillier_dec( + d_ct, d_exp_d, d_result, NUM_DEC, TAU, + d_r_0, + d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, + MOD, + d_con_twiddle, d_con_twiddle_shoup, + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, + d_con_InvTwiddle, d_con_InvTwiddle_shoup, + inv, inv_shoup, d_sample + ); + clock_gettime(CLOCK_MONOTONIC, &t1); + round_ms[r] = (double)(t1.tv_sec - t0.tv_sec ) * 1.0e3 + + (double)(t1.tv_nsec - t0.tv_nsec) * 1.0e-6; + } + + // ── 7. D2H result ───────────────────────────────────────────────────────── + CUDA_CHECK(cudaMemcpy(h_result, d_result, ct_bytes, cudaMemcpyDeviceToHost)); + + // ── 8. 写结果到 dec_results.txt ────────────────────────────────────────── + // 格式: + // 行1: NUM_DEC ARR_LEN BASE_BITS TAU + // 行2: n² limbs(ARR_LEN 个 uint64,空格分隔) + // 行3: r₀ = R mod n² limbs(ARR_LEN 个 uint64) + // 行4: λ limbs(EXP_U64_LIMBS 个 uint64) + // 每个用例 2 行:c̃ limbs / result limbs + { + FILE *fout = fopen("dec_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[Dec] 无法创建 dec_results.txt\n"); + } else { + fprintf(fout, "%d %d %d %d\n", NUM_DEC, ARR_LEN, BASE_BITS, TAU); + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 max_ms) max_ms = round_ms[r]; + } + double avg_ms = sum_ms / kRounds; + + printf("\n============================================================\n"); + printf(" paillier_dec 性能报告(非 CRT,GPU 批量模幂)\n"); + printf(" 批大小 NUM_DEC = %d,预热 %d 轮,计时 %d 轮\n", NUM_DEC, kWarmup, + kRounds); printf(" 指数 TAU = %d bit,大数 ARR_LEN = %d × %d bit\n", TAU, + ARR_LEN, BASE_BITS); + printf("============================================================\n"); + printf("\n[paillier_dec 计算耗时]\n"); + printf("------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %10.3f ms (%.3f us/次)\n", + r, round_ms[r], round_ms[r] * 1e3 / NUM_DEC); + printf("------------------------------------------------------------\n"); + printf(" 平均每轮总耗时 : %10.3f ms\n", avg_ms); + printf(" 最快轮 : %10.3f ms\n", min_ms); + printf(" 最慢轮 : %10.3f ms\n", max_ms); + printf(" 平均单次 dec : %10.3f us\n", avg_ms * 1e3 / NUM_DEC); + printf(" (最快轮) 单次 : %10.3f us\n", min_ms * 1e3 / NUM_DEC); + printf("============================================================\n"); + + // ── 10. 释放 paillier_dec 缓冲区 ───────────────────────────────────────── + cudaFreeHost(h_ct); + cudaFreeHost(h_exp); + cudaFreeHost(h_result); + cudaFree(d_ct); + cudaFree(d_exp_d); + cudaFree(d_result); + + // ── 11. 释放公共设备资源 ────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_dec 测试完成。\n"); + return 0; + */ + // ============================================================ + // [paillier_subhomo2 测试] + // 全 GPU 密文同态减法:c̃_result = c₁ · c₂⁻¹ · R (mod n²) + // + // 流程: + // ① CPU 生成随机 c₁, c₂ ∈ (0, n²)(标准域,约 4095 bit) + // ② GPU XYfixWarpROneVector:c₁,c₂ → c̃₁,c̃₂ (标准域 → FMLM 域) + // ③ paillier_subhomo2(c̃₁, c̃₂) → c̃_result (全 GPU,无 D2H/H2D) + // ④ 结果写入 subhomo2_results.txt,由 verify_subhomo2.py 验证 + // ============================================================ + { + const int NUM_SUB = 10000; // 测试用例对数 + const int kWarmup = 3; + const int kRounds = 2; + const size_t sub_bytes = (size_t)NUM_SUB * ARR_LEN * sizeof(uint64_t); + + printf("\n============================================================\n"); + printf(" [paillier_subhomo2] 全 GPU 密文同态减法测试\n"); + printf(" 批大小=%d,密文约 4095 bit(< n² ≈ 2^4095)\n", NUM_SUB); + printf("============================================================\n"); + + // ── 1. 计算 r_param = (2^4352 − 1) mod n²(FMLM 恒等元)──────────────── + // bn17_to_mp_u64 / mp_to_bn17_u64 与 paillier_subhomo 共用, + // 定义在本文件 SubHomo 辅助函数区(static)。 + uint64_t h_r_param[ARR_LEN] = {}; + { + mp_int r_big, n2_mp, r_param_mp; + mp_init(&r_big); + mp_init(&n2_mp); + mp_init(&r_param_mp); + mp_2expt(&r_big, ARR_LEN * BASE_BITS); // r_big = 2^4352 + mp_sub_d(&r_big, 1, &r_big); // r_big = 2^4352 - 1 + bn17_to_mp_u64(N2_arr, &n2_mp); + mp_mod(&r_big, &n2_mp, &r_param_mp); // r_param = r_big mod n² + mp_to_bn17_u64(&r_param_mp, h_r_param); + mp_clear(&r_big); + mp_clear(&n2_mp); + mp_clear(&r_param_mp); + printf("[SubHomo2] r_param(FMLM 恒等元 = (2^4352-1) mod n²)计算完成\n"); + } + + // ── 2. CPU 随机生成 c₁, c₂ ∈ (0, n²)(标准域,约 4095 bit)─────────── + // n² 最高非零 limb:index=N2_TOP_IDX=240,值=N2_TOP_VAL=27958 + // 策略:limb[0..239] 全随机 17-bit,limb[240] ∈ [1, + // N2_TOP_VAL),limb[241..] = 0 + uint64_t *h_c1_std = NULL; + uint64_t *h_c2_std = NULL; + uint64_t *h_result = NULL; + + h_c1_std = (uint64_t *)malloc(sub_bytes); + h_c2_std = (uint64_t *)malloc(sub_bytes); + h_result = (uint64_t *)malloc(sub_bytes); + + if (!h_c1_std || !h_c2_std || !h_result) { + fprintf(stderr, "[SubHomo2] host malloc 失败,跳过测试\n"); + free(h_c1_std); + free(h_c2_std); + free(h_result); + goto sub_done; + } + + srand((unsigned int)time(NULL)); + { + const uint64_t MASK17 = (1ULL << BASE_BITS) - 1; + for (int p = 0; p < NUM_SUB; p++) { + uint64_t *c1p = h_c1_std + (size_t)p * ARR_LEN; + uint64_t *c2p = h_c2_std + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < N2_TOP_IDX) { + // 0..239:随机 17-bit limb + c1p[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17; + c2p[j] = ((uint64_t)(uint32_t)rand() ^ + ((uint64_t)(uint32_t)rand() << 15)) & + MASK17; + } else if (j == N2_TOP_IDX) { + // limb 240:[1, N2_TOP_VAL),保证 c < n² 且高位非零(约 4095 bit) + c1p[j] = 1 + (uint64_t)(rand() % (int)(N2_TOP_VAL - 1)); + c2p[j] = 1 + (uint64_t)(rand() % (int)(N2_TOP_VAL - 1)); + } else { + c1p[j] = 0ULL; + c2p[j] = 0ULL; + } + } + } + } + printf("[SubHomo2] 随机 c₁, c₂ 生成完成(%d 对,各约 4095 bit)\n", + NUM_SUB); + + // ── 3. 分配 GPU 缓冲区 ─────────────────────────────────────────────────── + { + uint64_t *d_c1_std = NULL; // 标准域 c₁ batch + uint64_t *d_c2_std = NULL; // 标准域 c₂ batch + uint64_t *d_c1_tilde = NULL; // FMLM 域 c̃₁ + uint64_t *d_c2_tilde = NULL; // FMLM 域 c̃₂ + uint64_t *d_sub_res = NULL; // SubHomo2 结果 + uint64_t *d_ctR_bat = NULL; // R² CT 广播 NUM_SUB 份 + uint64_t *d_nctR_bat = NULL; // R² NCT 广播 NUM_SUB 份 + + CUDA_CHECK(cudaMalloc(&d_c1_std, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_std, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_c1_tilde, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_c2_tilde, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_sub_res, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_ctR_bat, sub_bytes)); + CUDA_CHECK(cudaMalloc(&d_nctR_bat, sub_bytes)); + + // H2D:标准域密文(计时:c₁ + c₂ 一起传输) + float ms_h2d = 0.f; + { + cudaEvent_t ev0, ev1; + cudaEventCreate(&ev0); + cudaEventCreate(&ev1); + cudaEventRecord(ev0); + CUDA_CHECK( + cudaMemcpy(d_c1_std, h_c1_std, sub_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_c2_std, h_c2_std, sub_bytes, cudaMemcpyHostToDevice)); + cudaEventRecord(ev1); + cudaEventSynchronize(ev1); + cudaEventElapsedTime(&ms_h2d, ev0, ev1); + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + printf( + "[SubHomo2] H2D(c₁+c₂ → GPU): %.3f ms" + " 数据量=%.2f MB 带宽=%.2f GB/s\n", + ms_h2d, 2.0 * sub_bytes / 1048576.0, + 2.0 * sub_bytes / 1e9 / (ms_h2d * 1e-3)); + } + + // ── 4. 广播 R²(CT/NCT)到 NUM_SUB 份 ────────────────────────────── + { + const int total = NUM_SUB * ARR_LEN; + const int blk = 256, grd = (total + blk - 1) / blk; + broadcast_fill_kernel<<>>(d_ctR_bat, d_ctR, NUM_SUB); + CUDA_CHECK(cudaGetLastError()); + broadcast_fill_kernel<<>>(d_nctR_bat, d_nctR, NUM_SUB); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + } + + // ── 5. 标准域 → FMLM 域:c̃₁ = c₁·R, c̃₂ = c₂·R mod n² ──────────── + // XYfixWarpROneVector 按 (blockIdx.x*32+threadIdx.x) 索引 inoutAct, + // 故 d_ctR_bat/d_nctR_bat 必须已广播为 batch×ARR_LEN。 + CUDA_CHECK(cudaMemcpy(d_c1_tilde, d_c1_std, sub_bytes, + cudaMemcpyDeviceToDevice)); + CUDA_CHECK(cudaMemcpy(d_c2_tilde, d_c2_std, sub_bytes, + cudaMemcpyDeviceToDevice)); + + XYfixWarpROneVector<<>>( + d_c1_tilde, d_ctR_bat, d_nctR_bat, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaGetLastError()); + XYfixWarpROneVector<<>>( + d_c2_tilde, d_ctR_bat, d_nctR_bat, d_negModn, d_con_NegModn_shoup, + d_Modn, d_con_Modn_shoup, (uint64_t)ARR_LEN, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaDeviceSynchronize()); + printf("[SubHomo2] c₁, c₂ 已转为 FMLM 域(c̃₁, c̃₂)\n"); + + // ── 6. 预热 ─────────────────────────────────────────────────────── + printf("[SubHomo2] 预热 %d 次...\n", kWarmup); + for (int w = 0; w < kWarmup; w++) { + paillier_subhomo2(NUM_SUB, d_c1_tilde, d_c2_tilde, d_sub_res, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, + d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, + d_ctR, d_nctR, N2_arr, false); + CUDA_CHECK(cudaDeviceSynchronize()); + } + printf("[SubHomo2] 预热完成\n\n"); + + // ── 7. cudaEvent 计时(kRounds 轮)────────────────────────────── + printf("[SubHomo2] 开始计时(%d 轮 × %d 对密文)...\n", kRounds, NUM_SUB); + float sub_ms[5] = {}; + { + cudaEvent_t ev_start, ev_stop; + cudaEventCreate(&ev_start); + cudaEventCreate(&ev_stop); + for (int rnd = 0; rnd < kRounds; rnd++) { + cudaEventRecord(ev_start); + paillier_subhomo2( + NUM_SUB, d_c1_tilde, d_c2_tilde, d_sub_res, d_negModn, + d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, + d_con_twiddle_shoup, d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, + d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, d_con_InvTwiddle, + d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample, d_ctR, d_nctR, + N2_arr, false); + cudaEventRecord(ev_stop); + cudaEventSynchronize(ev_stop); + cudaEventElapsedTime(&sub_ms[rnd], ev_start, ev_stop); + } + cudaEventDestroy(ev_start); + cudaEventDestroy(ev_stop); + } + printf("[SubHomo2] 计时完成\n\n"); + + // ── 8. D2H 最后一轮结果(计时)──────────────────────────────────── + float ms_d2h = 0.f; + { + cudaEvent_t ev0, ev1; + cudaEventCreate(&ev0); + cudaEventCreate(&ev1); + cudaEventRecord(ev0); + CUDA_CHECK( + cudaMemcpy(h_result, d_sub_res, sub_bytes, cudaMemcpyDeviceToHost)); + cudaEventRecord(ev1); + cudaEventSynchronize(ev1); + cudaEventElapsedTime(&ms_d2h, ev0, ev1); + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + printf( + "[SubHomo2] D2H(result → CPU): %.3f ms" + " 数据量=%.2f MB 带宽=%.2f GB/s\n", + ms_d2h, (double)sub_bytes / 1048576.0, + (double)sub_bytes / 1e9 / (ms_d2h * 1e-3)); + } + /* + // ── 9. 写入 subhomo2_results.txt ───────────────────────────────── + // 格式(供 verify_subhomo2.py 验证): + // 行1: NUM_SUB ARR_LEN BASE_BITS + // 行2: n² limbs(ARR_LEN 个 uint64,base-2^BASE_BITS,小端序) + // 行3: r_param limbs(FMLM 恒等元 = (2^4352-1) mod n²) + // 每个用例 3 行: + // c₁ limbs(标准域) + // c₂ limbs(标准域) + // result limbs(FMLM 域,= c₁·c₂⁻¹·R mod n²) + { + FILE *fout = fopen("subhomo2_results.txt", "w"); + if (!fout) { + fprintf(stderr, "[SubHomo2] 无法创建 subhomo2_results.txt\n"); + } else { + // 行1:元参数 + fprintf(fout, "%d %d %d\n", NUM_SUB, ARR_LEN, BASE_BITS); + // 行2:n² + for (int j = 0; j < ARR_LEN; j++) + fprintf(fout, "%llu%c", (unsigned long long)N2_arr[j], + j+1 rmax) rmax = sub_ms[r]; + } + float ravg = rsum / kRounds; + float e2e_avg = ms_h2d + ravg + ms_d2h; // 端到端平均(含传输) + + printf( + "\n============================================================\n"); + printf(" paillier_subhomo2 性能报告\n"); + printf(" 批大小=%d 密文约 4095 bit\n", NUM_SUB); + printf( + "============================================================\n"); + + // ── 数据传输耗时 ──────────────────────────────────────────────── + printf("[数据传输(cudaMemcpy,仅测一次)]\n"); + printf( + "------------------------------------------------------------\n"); + printf( + " H2D(c₁+c₂ → GPU) : %8.3f ms" + " (%.3f us/对 带宽=%.2f GB/s)\n", + ms_h2d, ms_h2d * 1e3f / NUM_SUB, + 2.0 * sub_bytes / 1e9 / (ms_h2d * 1e-3)); + printf( + " D2H(result → CPU) : %8.3f ms" + " (%.3f us/对 带宽=%.2f GB/s)\n", + ms_d2h, ms_d2h * 1e3f / NUM_SUB, + (double)sub_bytes / 1e9 / (ms_d2h * 1e-3)); + printf(" 传输小计 : %8.3f ms\n", ms_h2d + ms_d2h); + printf( + "------------------------------------------------------------\n"); + + // ── GPU 计算耗时 ──────────────────────────────────────────────── + printf("[GPU 计算(cudaEvent,%d 轮,不含传输)]\n", kRounds); + printf( + "------------------------------------------------------------\n"); + for (int r = 0; r < kRounds; r++) + printf(" 轮 %d : %8.3f ms (%.3f us/对)\n", r, sub_ms[r], + sub_ms[r] * 1e3f / NUM_SUB); + printf( + "------------------------------------------------------------\n"); + printf(" GPU 计算平均 : %8.3f ms (%.3f us/对)\n", ravg, + ravg * 1e3f / NUM_SUB); + printf(" GPU 计算最快 : %8.3f ms (%.3f us/对)\n", rmin, + rmin * 1e3f / NUM_SUB); + printf(" GPU 计算最慢 : %8.3f ms (%.3f us/对)\n", rmax, + rmax * 1e3f / NUM_SUB); + printf( + "------------------------------------------------------------\n"); + + // ── 端到端总耗时(H2D + GPU计算平均 + D2H)──────────────────── + printf("[端到端总耗时(H2D + GPU计算 + D2H)]\n"); + printf( + "------------------------------------------------------------\n"); + printf(" 端到端平均 : %8.3f ms (%.3f us/对)\n", e2e_avg, + e2e_avg * 1e3f / NUM_SUB); + printf(" 其中 H2D : %8.3f ms 占比 %.1f%%\n", ms_h2d, + ms_h2d / e2e_avg * 100.f); + printf(" 其中 GPU 计算 : %8.3f ms 占比 %.1f%%\n", ravg, + ravg / e2e_avg * 100.f); + printf(" 其中 D2H : %8.3f ms 占比 %.1f%%\n", ms_d2h, + ms_d2h / e2e_avg * 100.f); + printf( + "============================================================\n"); + } + + // ── 11. 释放 SubHomo2 GPU 缓冲区 ───────────────────────────────── + cudaFree(d_c1_std); + cudaFree(d_c2_std); + cudaFree(d_c1_tilde); + cudaFree(d_c2_tilde); + cudaFree(d_sub_res); + cudaFree(d_ctR_bat); + cudaFree(d_nctR_bat); + } + + free(h_c1_std); + free(h_c2_std); + free(h_result); + printf("[SubHomo2] 测试完成\n"); + sub_done:; + } + + // ── 释放公共设备资源 ────────────────────────────────────────────────────── + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_IRct); + cudaFree(d_IRnct); + cudaFree(d_r_0); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] paillier_subhomo2 测试完成。\n"); + return 0; + + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +} diff --git a/heu/library/algorithms/paillier_new/pallier_caddandsubp2.cu b/heu/library/algorithms/paillier_new/pallier_caddandsubp2.cu new file mode 100644 index 0000000..5c9beb3 --- /dev/null +++ b/heu/library/algorithms/paillier_new/pallier_caddandsubp2.cu @@ -0,0 +1,17235 @@ +// Copyright 2024 Ant Group Co., Ltd. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include +#include +#include +#include + +#include +// #include +/// #include +#include + +#include +// #include +#include + +#include +#include +// #include "heu/library/algorithms/paillier_cuda/gpupaillier/FMLMTest128.h" +#include +#include + +using namespace std; + +// const uint64_t MOD = 269670610521107713; +// const uint64_t G = 5; +// const ll P = 25668312996353, g = 3, gi = 8556104332118; +const uint64_t MOD = 25668312996353; +// const uint64_t G = 3; +const uint64_t inv = 25568046148711; +const uint64_t inv_shoup = 18374686479671626487; +const uint64_t inv_128 = 25467779301069; +const uint64_t inv_128_shoup = 18302628885633701358; +const uint64_t COUNT_streams = 10; + +//__uint128_t tempshoup=inv<<64; +// uint64_t inv_shoup=tempshoup/MOD; +/* +//生成随机元素 +uint64_t getRand(uint64_t min, uint64_t max) { + return (uint64_t)(rand() % (max - min + 1)) + min; +} +*/ + +const uint64_t con_twiddle[256] = {1, + 1, + 1, + 14042638072511, + 1, + 3231476038289, + 14042638072511, + 12668264113833, + 1, + 2922023234955, + 3231476038289, + 24491270218746, + 14042638072511, + 15090229196712, + 12668264113833, + 25398880845211, + 1, + 11659442688204, + 2922023234955, + 23079846529314, + 3231476038289, + 12783635489077, + 24491270218746, + 2828664205064, + 14042638072511, + 15141975243488, + 15090229196712, + 24633365557971, + 12668264113833, + 6510430045738, + 25398880845211, + 7958405596020, + 1, + 18677361289931, + 11659442688204, + 12238775288292, + 2922023234955, + 13580919637767, + 23079846529314, + 3355963859057, + 3231476038289, + 23225787056075, + 12783635489077, + 19331914697329, + 24491270218746, + 4037356296035, + 2828664205064, + 2411517338904, + 14042638072511, + 22062777184955, + 15141975243488, + 7322853327189, + 15090229196712, + 20471538997241, + 24633365557971, + 11333564164083, + 12668264113833, + 15566904200236, + 6510430045738, + 20284254982169, + 25398880845211, + 25632652396295, + 7958405596020, + 12611594884277, + 1, + 14547619135716, + 18677361289931, + 11334678972964, + 11659442688204, + 20041366642069, + 12238775288292, + 3793617661074, + 2922023234955, + 1765551122505, + 13580919637767, + 20104284704739, + 23079846529314, + 4647857287623, + 3355963859057, + 22599399742732, + 3231476038289, + 5275640180651, + 23225787056075, + 5256631699163, + 12783635489077, + 13884774419207, + 19331914697329, + 8570623566421, + 24491270218746, + 13975252702604, + 4037356296035, + 10547486611031, + 2828664205064, + 3804246056013, + 2411517338904, + 24647895302756, + 14042638072511, + 22359154462675, + 22062777184955, + 3213717192640, + 15141975243488, + 1163936124732, + 7322853327189, + 22981887500849, + 15090229196712, + 21905340849536, + 20471538997241, + 15332924377786, + 24633365557971, + 950604996099, + 11333564164083, + 20391393316196, + 12668264113833, + 22307962513487, + 15566904200236, + 2749801639527, + 6510430045738, + 12550400983564, + 20284254982169, + 19119151774540, + 25398880845211, + 5262440844454, + 25632652396295, + 104582330697, + 7958405596020, + 6412318858354, + 12611594884277, + 15745789325549, + 1, + 9686405673677, + 14547619135716, + 25276323353834, + 18677361289931, + 15273987809805, + 11334678972964, + 16453600334917, + 11659442688204, + 18778247373441, + 20041366642069, + 4983427385556, + 12238775288292, + 11952650055719, + 3793617661074, + 8366594500092, + 2922023234955, + 4226361580066, + 1765551122505, + 17094862308505, + 13580919637767, + 23150412329133, + 20104284704739, + 3453012878703, + 23079846529314, + 18616498543012, + 4647857287623, + 23335983881574, + 3355963859057, + 5862003252156, + 22599399742732, + 10340544793251, + 3231476038289, + 8090678643853, + 5275640180651, + 17787453245469, + 23225787056075, + 20899621539132, + 5256631699163, + 15757637799910, + 12783635489077, + 3291408355812, + 13884774419207, + 6044007603010, + 19331914697329, + 7251131673046, + 8570623566421, + 4242187245377, + 24491270218746, + 15212546819995, + 13975252702604, + 1173099840187, + 4037356296035, + 22960681611827, + 10547486611031, + 13127743180518, + 2828664205064, + 20126412828462, + 3804246056013, + 7320053948647, + 2411517338904, + 22440872643659, + 24647895302756, + 20247606826077, + 14042638072511, + 12670157626148, + 22359154462675, + 18988571635575, + 22062777184955, + 21946518500003, + 3213717192640, + 18588479514129, + 15141975243488, + 25634161461568, + 1163936124732, + 24589958172766, + 7322853327189, + 4479952438867, + 22981887500849, + 18737494629827, + 15090229196712, + 536834336436, + 21905340849536, + 17860342618587, + 20471538997241, + 19686162102092, + 15332924377786, + 753843377339, + 24633365557971, + 12740557981160, + 950604996099, + 2158098352461, + 11333564164083, + 19580334808271, + 20391393316196, + 22031342610290, + 12668264113833, + 13645228754854, + 22307962513487, + 3800482705353, + 15566904200236, + 15713377562347, + 2749801639527, + 14754382421516, + 6510430045738, + 5915545648504, + 12550400983564, + 24321837644922, + 20284254982169, + 17510182993527, + 19119151774540, + 7621823278492, + 25398880845211, + 18984727940996, + 5262440844454, + 17444676403103, + 25632652396295, + 24264202146457, + 104582330697, + 12140592244261, + 7958405596020, + 21443156625714, + 6412318858354, + 7554505135329, + 12611594884277, + 7259027004882, + 15745789325549, + 16036524428840}; +const uint64_t con_twiddle_shoup[256] = {718658, + 718658, + 718658, + 10091857251395735022, + 718658, + 2322326810769042622, + 10091857251395735022, + 9104152111564134350, + 718658, + 2099936010609899669, + 2322326810769042622, + 17600852608796780638, + 10091857251395735022, + 10844717221771318485, + 9104152111564134350, + 18253114444132739460, + 718658, + 8379154303671693428, + 2099936010609899669, + 16586521375488916475, + 2322326810769042622, + 9187064698490293466, + 17600852608796780638, + 2032842776566284642, + 10091857251395735022, + 10881904943529218931, + 10844717221771318485, + 17702970592051969299, + 9104152111564134350, + 4678774054242736776, + 18253114444132739460, + 5719373582728924273, + 718658, + 13422639179151322714, + 8379154303671693428, + 8795496437605406151, + 2099936010609899669, + 9760039503924505741, + 16586521375488916475, + 2411791006188796026, + 2322326810769042622, + 16691402734366106000, + 9187064698490293466, + 13893039364413244670, + 17600852608796780638, + 2901479280618063106, + 2032842776566284642, + 1733056753141339998, + 10091857251395735022, + 15855596133020830364, + 10881904943529218931, + 5262628721847079115, + 10844717221771318485, + 14712039732830287778, + 17702970592051969299, + 8144959024284521310, + 9104152111564134350, + 11187283630307444966, + 4678774054242736776, + 14577446536322951759, + 18253114444132739460, + 18421116290423665141, + 5719373582728924273, + 9063426304043306776, + 718658, + 10454766042337154141, + 13422639179151322714, + 8145760190848207337, + 8379154303671693428, + 14402892834685083217, + 8795496437605406151, + 2726314507590723584, + 2099936010609899669, + 1268827823259275760, + 9760039503924505741, + 14448109417475390221, + 16586521375488916475, + 3340220835241065961, + 2411791006188796026, + 16241244344062132348, + 2322326810769042622, + 3791382170354358438, + 16691402734366106000, + 3777721568923761814, + 9187064698490293466, + 9978407239645704801, + 13893039364413244670, + 6159349058285799670, + 17600852608796780638, + 10043430201547804164, + 2901479280618063106, + 7580037930899913062, + 2032842776566284642, + 2733952690956412607, + 1733056753141339998, + 17713412512545275294, + 10091857251395735022, + 16068590099244491768, + 15855596133020830364, + 2309564270403489249, + 10881904943529218931, + 836472261115371968, + 5262628721847079115, + 16516122314667902277, + 10844717221771318485, + 15742453216780498573, + 14712039732830287778, + 11019132108087796203, + 17702970592051969299, + 683160092395608723, + 8144959024284521310, + 14654442380520444617, + 9104152111564134350, + 16031800584271709924, + 11187283630307444966, + 1976167545760743178, + 4678774054242736776, + 9019448804412338064, + 14577446536322951759, + 13740135541492714252, + 18253114444132739460, + 3781896358925944562, + 18421116290423665141, + 75158951399483166, + 5719373582728924273, + 4608265643164160887, + 9063426304043306776, + 11315840895668484533, + 718658, + 6961215039018548975, + 10454766042337154141, + 18165037495793025788, + 13422639179151322714, + 10976776859175148089, + 8145760190848207337, + 11824515094250240362, + 8379154303671693428, + 13495141792134586415, + 14402892834685083217, + 3581381043792333544, + 8795496437605406151, + 8589870187876618994, + 2726314507590723584, + 6012721893083952433, + 2099936010609899669, + 3037309481207950433, + 1268827823259275760, + 12285363281377319705, + 9760039503924505741, + 16637234067429468313, + 14448109417475390221, + 2481536081693700950, + 16586521375488916475, + 13378899665915791204, + 3340220835241065961, + 16770596588594239490, + 2411791006188796026, + 4212776810347404146, + 16241244344062132348, + 7431317492931189510, + 2322326810769042622, + 5814432695556930281, + 3791382170354358438, + 12783099449849351990, + 16691402734366106000, + 15019684769487286100, + 3777721568923761814, + 11324355899137248940, + 9187064698490293466, + 2365397663272990145, + 9978407239645704801, + 4343575732776875929, + 13893039364413244670, + 5211085365697923150, + 6159349058285799670, + 3048682725636973809, + 17600852608796780638, + 10932621786934074045, + 10043430201547804164, + 843057840533337593, + 2901479280618063106, + 16500882528254985291, + 7580037930899913062, + 9434360518776925845, + 2032842776566284642, + 14464011975434986499, + 2733952690956412607, + 5260616925452939570, + 1733056753141339998, + 16127317541558098648, + 17713412512545275294, + 14551109037777623266, + 10091857251395735022, + 9105512899749948708, + 16068590099244491768, + 13646293051534802838, + 15855596133020830364, + 15772045873681220487, + 2309564270403489249, + 13358763560552263286, + 10881904943529218931, + 18422200792583410069, + 836472261115371968, + 17671775517958111637, + 5262628721847079115, + 3219554635862208799, + 16516122314667902277, + 13465854498035602581, + 10844717221771318485, + 385800407514961996, + 15742453216780498573, + 12835481996836055400, + 14712039732830287778, + 14147622173005586619, + 11019132108087796203, + 541755738112168677, + 17702970592051969299, + 9156106693420345944, + 683160092395608723, + 1550935115969583168, + 8144959024284521310, + 14071568518626094719, + 14654442380520444617, + 15833005417612554117, + 9104152111564134350, + 9806255779403095303, + 16031800584271709924, + 2731248128077874018, + 11187283630307444966, + 11292547915687684068, + 1976167545760743178, + 10603358176831297944, + 4678774054242736776, + 4251255493487391129, + 9019448804412338064, + 17479088497243158013, + 14577446536322951759, + 12583836904720119956, + 13740135541492714252, + 5477485934247420062, + 18253114444132739460, + 13643530748838618951, + 3781896358925944562, + 12536760055187214146, + 18421116290423665141, + 17437668272630204514, + 75158951399483166, + 8724936386156955069, + 5719373582728924273, + 15410300726160570119, + 4608265643164160887, + 5429107197451525502, + 9063426304043306776, + 5216759410804681689, + 11315840895668484533, + 11524780066872084464}; +const uint64_t con_twiddle_NCT[256] = {1, + 14042638072511, + 3231476038289, + 12668264113833, + 2922023234955, + 15090229196712, + 24491270218746, + 25398880845211, + 11659442688204, + 15141975243488, + 12783635489077, + 6510430045738, + 23079846529314, + 24633365557971, + 2828664205064, + 7958405596020, + 18677361289931, + 22062777184955, + 23225787056075, + 15566904200236, + 13580919637767, + 20471538997241, + 4037356296035, + 25632652396295, + 12238775288292, + 7322853327189, + 19331914697329, + 20284254982169, + 3355963859057, + 11333564164083, + 2411517338904, + 12611594884277, + 14547619135716, + 22359154462675, + 5275640180651, + 22307962513487, + 1765551122505, + 21905340849536, + 13975252702604, + 5262440844454, + 20041366642069, + 1163936124732, + 13884774419207, + 12550400983564, + 4647857287623, + 950604996099, + 3804246056013, + 6412318858354, + 11334678972964, + 3213717192640, + 5256631699163, + 2749801639527, + 20104284704739, + 15332924377786, + 10547486611031, + 104582330697, + 3793617661074, + 22981887500849, + 8570623566421, + 19119151774540, + 22599399742732, + 20391393316196, + 24647895302756, + 15745789325549, + 9686405673677, + 12670157626148, + 8090678643853, + 13645228754854, + 4226361580066, + 536834336436, + 15212546819995, + 18984727940996, + 18778247373441, + 25634161461568, + 3291408355812, + 5915545648504, + 18616498543012, + 12740557981160, + 20126412828462, + 21443156625714, + 15273987809805, + 21946518500003, + 20899621539132, + 15713377562347, + 23150412329133, + 19686162102092, + 22960681611827, + 24264202146457, + 11952650055719, + 4479952438867, + 7251131673046, + 17510182993527, + 5862003252156, + 19580334808271, + 22440872643659, + 7259027004882, + 25276323353834, + 18988571635575, + 17787453245469, + 3800482705353, + 17094862308505, + 17860342618587, + 1173099840187, + 17444676403103, + 4983427385556, + 24589958172766, + 6044007603010, + 24321837644922, + 23335983881574, + 2158098352461, + 7320053948647, + 7554505135329, + 16453600334917, + 18588479514129, + 15757637799910, + 14754382421516, + 3453012878703, + 753843377339, + 13127743180518, + 12140592244261, + 8366594500092, + 18737494629827, + 4242187245377, + 7621823278492, + 10340544793251, + 22031342610290, + 20247606826077, + 16036524428840, + 24428787703014, + 12846003271633, + 19975165085819, + 3417089886880, + 25380161891250, + 58309668949, + 13704088879039, + 8106971142527, + 22437342183330, + 12345120229153, + 7132141190150, + 24311014668732, + 6117259658301, + 4392921321090, + 5248675351858, + 13491483347253, + 5910204534378, + 13139550084763, + 11165709197323, + 14097072216406, + 22864570519176, + 13518431242485, + 17022042066798, + 17924644593086, + 22194371404954, + 9423851108789, + 22740835540735, + 3196497611009, + 8026020327984, + 7553319980811, + 5187517429973, + 17659663135745, + 16411136245780, + 17826102210926, + 11111936218880, + 15779912951768, + 16561767585736, + 14975726082350, + 2432723066402, + 18964612531886, + 15757409314132, + 7048612611320, + 4390575991869, + 21270379752021, + 13498830613605, + 19591405666399, + 19258601875796, + 18090581333047, + 16535903207636, + 6997440895456, + 7550050743936, + 4157340647969, + 10619772370371, + 9779370495189, + 21879658639537, + 20340744241821, + 11273398814987, + 5706105868532, + 1044437897516, + 18992007532740, + 22023149280814, + 7032268118645, + 18276483659079, + 24017782760760, + 22133570538252, + 1640114928618, + 9014397249816, + 16576775818355, + 4012792011599, + 15834542960286, + 24865093073350, + 23624657739462, + 19076535165349, + 2903431988357, + 18701383276484, + 21301885339380, + 24069964557985, + 24291533392014, + 5045422665816, + 4309377652335, + 17081294988659, + 4135185860842, + 24649182537198, + 16801253268288, + 13936436739909, + 21394991642555, + 9962633193554, + 526826856424, + 12810190436384, + 10545560915699, + 2226579679469, + 14407639768526, + 6515615151855, + 14193396639806, + 6788627509161, + 11279058929908, + 15026572684036, + 13253857283968, + 20757625930366, + 15650867790444, + 12233872979945, + 23911547317294, + 4322912363116, + 24510956465584, + 9675611750387, + 17347205310068, + 12258560502910, + 7495046644715, + 11072004696589, + 8118092776260, + 7965558935453, + 14469046375094, + 24674151278387, + 15809638743405, + 2260748213614, + 17561238915161, + 9424222348605, + 7853528769485, + 21022907714497, + 25386893002310, + 21376993549979, + 13673563000266, + 9281381453752, + 22504174008283, + 5729539430945, + 4491754775632, + 7747526086493, + 19343958507522}; +const uint64_t con_twiddle_NCT_shoup[256] = {718658, + 10091857251395735022, + 2322326810769042622, + 9104152111564134350, + 2099936010609899669, + 10844717221771318485, + 17600852608796780638, + 18253114444132739460, + 8379154303671693428, + 10881904943529218931, + 9187064698490293466, + 4678774054242736776, + 16586521375488916475, + 17702970592051969299, + 2032842776566284642, + 5719373582728924273, + 13422639179151322714, + 15855596133020830364, + 16691402734366106000, + 11187283630307444966, + 9760039503924505741, + 14712039732830287778, + 2901479280618063106, + 18421116290423665141, + 8795496437605406151, + 5262628721847079115, + 13893039364413244670, + 14577446536322951759, + 2411791006188796026, + 8144959024284521310, + 1733056753141339998, + 9063426304043306776, + 10454766042337154141, + 16068590099244491768, + 3791382170354358438, + 16031800584271709924, + 1268827823259275760, + 15742453216780498573, + 10043430201547804164, + 3781896358925944562, + 14402892834685083217, + 836472261115371968, + 9978407239645704801, + 9019448804412338064, + 3340220835241065961, + 683160092395608723, + 2733952690956412607, + 4608265643164160887, + 8145760190848207337, + 2309564270403489249, + 3777721568923761814, + 1976167545760743178, + 14448109417475390221, + 11019132108087796203, + 7580037930899913062, + 75158951399483166, + 2726314507590723584, + 16516122314667902277, + 6159349058285799670, + 13740135541492714252, + 16241244344062132348, + 14654442380520444617, + 17713412512545275294, + 11315840895668484533, + 6961215039018548975, + 9105512899749948708, + 5814432695556930281, + 9806255779403095303, + 3037309481207950433, + 385800407514961996, + 10932621786934074045, + 13643530748838618951, + 13495141792134586415, + 18422200792583410069, + 2365397663272990145, + 4251255493487391129, + 13378899665915791204, + 9156106693420345944, + 14464011975434986499, + 15410300726160570119, + 10976776859175148089, + 15772045873681220487, + 15019684769487286100, + 11292547915687684068, + 16637234067429468313, + 14147622173005586619, + 16500882528254985291, + 17437668272630204514, + 8589870187876618994, + 3219554635862208799, + 5211085365697923150, + 12583836904720119956, + 4212776810347404146, + 14071568518626094719, + 16127317541558098648, + 5216759410804681689, + 18165037495793025788, + 13646293051534802838, + 12783099449849351990, + 2731248128077874018, + 12285363281377319705, + 12835481996836055400, + 843057840533337593, + 12536760055187214146, + 3581381043792333544, + 17671775517958111637, + 4343575732776875929, + 17479088497243158013, + 16770596588594239490, + 1550935115969583168, + 5260616925452939570, + 5429107197451525502, + 11824515094250240362, + 13358763560552263286, + 11324355899137248940, + 10603358176831297944, + 2481536081693700950, + 541755738112168677, + 9434360518776925845, + 8724936386156955069, + 6012721893083952433, + 13465854498035602581, + 3048682725636973809, + 5477485934247420062, + 7431317492931189510, + 15833005417612554117, + 14551109037777623266, + 11524780066872084464, + 17555949035392735808, + 9231885817954543477, + 14355316542250183122, + 2455719728409595596, + 18239661914038721202, + 41904722771526049, + 9848556091353542777, + 5826141433618649880, + 16124780347229795866, + 8871922103284150558, + 5125571877316295053, + 17471310476462394255, + 4396218924287681036, + 3157009007861284209, + 3772003674548789182, + 9695765378770737240, + 4247417057928155287, + 9442845647540209009, + 8024328673012013348, + 10130976796238639703, + 16431811501692800004, + 9715131705133536029, + 12233030416247244451, + 12881693139235081896, + 15950167400647412821, + 6772528043322289334, + 16342888342598106143, + 2297189276556073616, + 5767965465503985937, + 5428255476417718067, + 3728052031398232031, + 12691262037333585123, + 11793997927627536135, + 12810874846487230638, + 7985684280154560744, + 11340364120074321189, + 11902250377955809305, + 10762428617663264439, + 1748296423473707001, + 13629074644775883818, + 11324191695955222631, + 5065543377713044634, + 3155323517742999982, + 15286133206023066325, + 9701045552113487829, + 14079524681798663923, + 13840352502743257136, + 13000944941054748463, + 11883662730084126142, + 5028768403584830233, + 5425906012471332986, + 2987707021151337274, + 7631986685884636717, + 7028024971971042116, + 15723996485511632982, + 14618043006993300310, + 8101720701725613972, + 4100739874461456160, + 750593878105598638, + 13648762287268189153, + 15827117214052903467, + 5053797273735226401, + 13134545175378946433, + 17260576956056375333, + 15906472358138351991, + 1178682071703945222, + 6478270662732421549, + 11913036167663170611, + 2883825955742015333, + 11379624424638134991, + 17869503475285875347, + 16978054428849925014, + 13709508765079812823, + 2086575258461783974, + 13439902777197218587, + 15308774955287458818, + 17298077833455390530, + 17457310096854388004, + 3625934461420387405, + 3096969663759425487, + 12275613015475509886, + 2971785301317339986, + 17714337594155344959, + 12074358731784207460, + 10015534790972475086, + 15375686565199744939, + 7159728216180231345, + 378608449764160186, + 9206148629597595717, + 7578654014124455778, + 1600149784394581560, + 10354168719773703080, + 4682500373364251103, + 10200201134696718379, + 4878702947522493637, + 8105788389827683891, + 10798969945815712822, + 9524993455610679552, + 14917638458346029819, + 11247624754412362396, + 8791973353425297867, + 17184229981580383123, + 3106696494888695972, + 17615000291865942857, + 6953457897343123196, + 12466711653172759783, + 8809715244682969802, + 5386376864549579160, + 7956987163508041610, + 5834134087097019697, + 5724514388897669142, + 10398299082215356782, + 17732281585200709532, + 11361726804516967046, + 1624705282250485652, + 12620528662373362833, + 6772794837866858587, + 5644002989475333651, + 15108285392563206857, + 18244499282317391394, + 15362752088064002675, + 9826618417715597788, + 6670141054932727883, + 16172809587447573697, + 4117580596663547461, + 3228036482168455002, + 5567823290226777049, + 13901694747582920512}; +const uint64_t con_InvTwiddle[256] = {1, + 11625674923842, + 13000048882520, + 22436836958064, + 269432151142, + 1177042777607, + 10578083799641, + 22746289761398, + 17709907400333, + 22839648791289, + 1034947438382, + 2588466467039, + 19157882950615, + 12884677507276, + 10526337752865, + 14008870308149, + 13056718112076, + 23256795657449, + 14334748832270, + 22312349137296, + 5384058014184, + 6336398299024, + 18345459669164, + 13429537708061, + 35660600058, + 21630956700318, + 5196773999112, + 12087393358586, + 10101408796117, + 2442525940278, + 3605535811398, + 6990951706422, + 9922523670804, + 1020417693597, + 5276919680157, + 3068913253621, + 6549161221813, + 17097689429932, + 2686425495504, + 21874695335279, + 25563730665656, + 15120826385322, + 10335388618567, + 5564028291614, + 22918511356826, + 20411681297190, + 22454595803713, + 14333634023389, + 19255994137999, + 21864066940340, + 24717708000254, + 21020455708730, + 13117912012789, + 11783538577146, + 24504376871621, + 5626946354284, + 20405872151899, + 11693060293749, + 3762972146817, + 23902761873848, + 3360350482866, + 20392672815702, + 3309158533678, + 11120693860637, + 9631788567513, + 5420706170276, + 3636970386063, + 15327768203102, + 18046489717861, + 21426125750976, + 6930818366526, + 17301718496261, + 13527720752092, + 12540569815835, + 24914469619014, + 22215300117650, + 10913930574837, + 9910675196443, + 7079833482224, + 9214712661436, + 18113807861024, + 18348259047706, + 23510214643892, + 2332329114779, + 1346475351431, + 19624305393343, + 1078354823587, + 20684885610797, + 8223636593250, + 24495213156166, + 7807970377766, + 8573450687848, + 21867830291000, + 7880859750884, + 6679741360778, + 391989642519, + 18409285991471, + 3227440352694, + 6087978188082, + 19806309744197, + 8158130002826, + 18417181323307, + 21188360557486, + 13715662940634, + 1404110849896, + 2707631384526, + 5982150894261, + 2517900667220, + 9954935434006, + 4768691457221, + 3721794496350, + 10394325186548, + 4225156370639, + 5541900167891, + 12927755015193, + 7051814453341, + 19752767347849, + 22376904640541, + 34151534785, + 6890065622912, + 6683585055357, + 10455766176358, + 25131478659917, + 21441951416287, + 12023084241499, + 17577634352500, + 12998155370205, + 15981907322676, + 6324354488831, + 17920786909860, + 21176558220721, + 19938773565408, + 3164138988070, + 16386931542601, + 11994749996087, + 4291319446374, + 281419994043, + 4645405281856, + 17814784226868, + 16244090647748, + 8107074081192, + 23407564782739, + 9858674252948, + 994161717966, + 11199266621259, + 17702754060900, + 17550220220093, + 14596308299764, + 18173266351638, + 13409752493443, + 8321107686285, + 15992701245966, + 1157356530769, + 21345400633237, + 1756765679059, + 13434440016408, + 10017445205909, + 4910687065987, + 12414455712385, + 10641740312317, + 14389254066445, + 18879685487192, + 11474916356547, + 19152697844498, + 11260673227827, + 23441733316884, + 15122752080654, + 12858122559969, + 25141486139929, + 15705679802799, + 4273321353798, + 11731876256444, + 8867059728065, + 1019130459155, + 21533127135511, + 8587018007694, + 21358935344018, + 20622890330537, + 1376779604339, + 1598348438368, + 4366427656973, + 6966929719869, + 22764881007996, + 6591777831004, + 2043655256891, + 803219923003, + 9833770036067, + 21655520984754, + 9091537177998, + 16653915746537, + 24028198067735, + 3534742458101, + 1650530235593, + 7391829337274, + 18636044877708, + 3645163715539, + 6676305463613, + 24623875098837, + 19962207127821, + 14394914181366, + 5327568754532, + 3788654356816, + 15888942501164, + 15048540625982, + 21510972348384, + 18118262252417, + 18670872100897, + 9132409788717, + 7577731663306, + 6409711120557, + 6076907329954, + 12169482382748, + 4397933244332, + 21277737004484, + 18619700385033, + 9910903682221, + 6703700464467, + 23235589929951, + 10692586914003, + 9106545410617, + 9888400044585, + 14556376777473, + 7842210785427, + 9257176750573, + 8008649860608, + 20480795566380, + 18114993015542, + 17642292668369, + 22471815385344, + 2927477455618, + 16244461887564, + 3473941591399, + 7743668403267, + 8646270929555, + 12149881753868, + 2803742477177, + 11571240779947, + 14502603799030, + 12528762911590, + 19758108461975, + 12176829649100, + 20419637644495, + 21275391675263, + 19551053338052, + 1357298327621, + 18536171806203, + 13323192767200, + 3230970813023, + 17561341853826, + 11964224117314, + 25610003327404, + 288151105103, + 22251223109473, + 5693147910534, + 12822309724720, + 1239525293339}; +const uint64_t con_InvTwiddle_shoup[256] = {718658, + 8354886822313816593, + 9342591962145417265, + 16124417262940508993, + 193629629576812155, + 845891464912770977, + 7602026851938233130, + 16346808063099651946, + 12727370490980627342, + 16413901297143266973, + 743773481657582316, + 1860222698220635140, + 13767970019466814839, + 9259679375219258149, + 7564839130180332684, + 10067589770037858187, + 9383317769666244839, + 16713687320568211617, + 10301785049425030305, + 16034953067520755589, + 3869297537386599856, + 4553704709296306945, + 13184115351862472500, + 9651247636104145464, + 25627783285886474, + 15545264793091488509, + 3734704340879263837, + 8686704569785045874, + 7259460443402106649, + 1755341339343445615, + 2591147940688721251, + 5024104894558228901, + 7130903178041067082, + 733331561164276321, + 3792301693189106998, + 2205499729647419267, + 4706608532216837363, + 12287395015423751945, + 1930621759041649338, + 15720429566118828031, + 18371585122310068449, + 10866706142809638553, + 7427611965621755412, + 3998634656234161394, + 16470576527948808437, + 14669022504785789801, + 16137179803306062366, + 10300983882861344278, + 13838478430545390728, + 15712791382753139008, + 17763583981313942892, + 15106523238468485654, + 9427295269297213551, + 8468336834063846814, + 17610271812594179647, + 4043851239024468398, + 14664847714783607053, + 8403313872161747451, + 2704290856929053042, + 17177916250450275855, + 2414943489437841691, + 14655361903355193177, + 2378153974465059847, + 7991978031372397474, + 6921964006837467151, + 3895635035931928349, + 2613738656096997498, + 11015426580778362105, + 12969258139462131553, + 15398061348072577806, + 4980889575673949034, + 12434022180625599182, + 9721807687552596546, + 9012383554932625770, + 17904988335597382938, + 15965207992015850665, + 7843385896878253671, + 7122388174572302675, + 5087980513157288329, + 6622228979459311253, + 13017636876258026113, + 13186127148256612045, + 16895808957739968447, + 1676147485115312125, + 967655576466393602, + 14103168340932675686, + 774968555751439978, + 14865363029917218071, + 5909984018522337469, + 17603686233176214022, + 5611262076873496215, + 6161380792332231910, + 15715495945631677597, + 5663644623860199625, + 4800451022174748777, + 281706577916525827, + 13229984662904869926, + 2319426532151452967, + 4375175555083456896, + 14233967263362147469, + 5862907168989431659, + 13235658708011628465, + 15227189437847342816, + 9856873885832932621, + 1009075801079347101, + 1945861545454566324, + 4299121900703964996, + 1809510006280083302, + 7154196158021867547, + 3427059304222265515, + 2674698200028331128, + 7469967214534403526, + 3036443347548981496, + 3982732098274565116, + 9290637380289205671, + 5067844407793760411, + 14195488580222160486, + 16081346410436561470, + 24543281126141546, + 4951602281574965200, + 4803213324870932664, + 7514122286775477570, + 18060943666194589619, + 15409434592501601182, + 8640488294306456312, + 12632311378152621334, + 9341231173959602907, + 11485529034691002640, + 4545049326126631103, + 12878920783482774566, + 15218707591541096613, + 14329163477046004154, + 2273934486261977918, + 11776603018776823732, + 8620125655993953827, + 3083991985645548940, + 202244791392160221, + 3338458681146344758, + 12802741084234217964, + 11673949235842693028, + 5826215411336188782, + 16822038791459065963, + 7085017269192584569, + 714462488508842083, + 8048444991494194833, + 12722229684811882473, + 12612609986612531918, + 10489756910201510005, + 13060367209159972455, + 9637028829026581813, + 5980032420536791832, + 11493286176366428419, + 831743781843608758, + 15340047578820855643, + 1262514092129168492, + 9654770720284253748, + 7199119319297189219, + 3529105615363521796, + 8921750618098872063, + 7647774127893838793, + 10340955683881867724, + 13568041126187057978, + 8246542939012833236, + 13764243700345300512, + 8092575353935848535, + 16846594289314970055, + 10868090059585095837, + 9240595444111955898, + 18068135623945391429, + 11287015857529320270, + 3071057508509806676, + 8431209282737076529, + 6372385341925344155, + 732406479554206656, + 15474958772392211629, + 6171131058234041729, + 15349774409950126128, + 14820809612289164210, + 989433976855163611, + 1148666240254161085, + 3137969118422092797, + 5006841296512333028, + 16360168815247767641, + 4737235308629738792, + 1468689644859626601, + 577240598423676268, + 7067119649071416624, + 15562918117967536282, + 6533707906046381004, + 11968473410977130066, + 17268062002005606393, + 2540271715571199624, + 1186167117653176282, + 5312198898330605182, + 13392946799974325214, + 2619626859656648148, + 4797981786441362462, + 17696150195603952977, + 14346004199248095455, + 10345023371983937643, + 3828701066716251305, + 2722747588197918633, + 11418719101738509499, + 10814757387824914898, + 15459037052558214341, + 13020838061238218629, + 13417975670124721382, + 6563081343625425473, + 5445799132654803152, + 4606391570966294479, + 4367219391910887692, + 8745698521596063786, + 3160610867686485290, + 15291420555966551633, + 13381200695996506981, + 7122552377754328984, + 4817669428933667797, + 16698447650235844614, + 7684315456046287176, + 6544493695753742310, + 7106379953635230426, + 10461059793554990871, + 5635869227222320977, + 6652746146082015480, + 5755482036375966492, + 14718692042311319584, + 13018488597291833548, + 12678778608205565678, + 16149554797153477999, + 2103855731111445472, + 11674216030387262281, + 2496576673062138794, + 5565050934474469719, + 6213713657462307164, + 8731612368576015586, + 2014932572016751611, + 8315767277470911912, + 10422415400697538267, + 9003898426169342606, + 14199327015781396328, + 8750978694938814375, + 14674740399160762433, + 15289735065848267406, + 14050525149421870579, + 975433597247157360, + 13321172196393256562, + 9574821970425401057, + 2321963726479755749, + 12620602640090901735, + 8598187982356008838, + 18404839350938025566, + 207082159670830413, + 15991024345299956019, + 4091427531459368493, + 9214858255755008138, + 890795038316815807}; +const uint64_t con_ICTTwiddle[256] = {1, + 1, + 1, + 11625674923842, + 1, + 13000048882520, + 11625674923842, + 22436836958064, + 1, + 269432151142, + 13000048882520, + 10578083799641, + 11625674923842, + 1177042777607, + 22436836958064, + 22746289761398, + 1, + 17709907400333, + 269432151142, + 19157882950615, + 13000048882520, + 1034947438382, + 10578083799641, + 10526337752865, + 11625674923842, + 22839648791289, + 1177042777607, + 12884677507276, + 22436836958064, + 2588466467039, + 22746289761398, + 14008870308149, + 1, + 13056718112076, + 17709907400333, + 35660600058, + 269432151142, + 5384058014184, + 19157882950615, + 10101408796117, + 13000048882520, + 14334748832270, + 1034947438382, + 5196773999112, + 10578083799641, + 18345459669164, + 10526337752865, + 3605535811398, + 11625674923842, + 23256795657449, + 22839648791289, + 21630956700318, + 1177042777607, + 6336398299024, + 12884677507276, + 2442525940278, + 22436836958064, + 22312349137296, + 2588466467039, + 12087393358586, + 22746289761398, + 13429537708061, + 14008870308149, + 6990951706422, + 1, + 9922523670804, + 13056718112076, + 19255994137999, + 17709907400333, + 25563730665656, + 35660600058, + 20405872151899, + 269432151142, + 6549161221813, + 5384058014184, + 13117912012789, + 19157882950615, + 22918511356826, + 10101408796117, + 3360350482866, + 13000048882520, + 5276919680157, + 14334748832270, + 24717708000254, + 1034947438382, + 10335388618567, + 5196773999112, + 3762972146817, + 10578083799641, + 2686425495504, + 18345459669164, + 24504376871621, + 10526337752865, + 22454595803713, + 3605535811398, + 3309158533678, + 11625674923842, + 1020417693597, + 23256795657449, + 21864066940340, + 22839648791289, + 15120826385322, + 21630956700318, + 11693060293749, + 1177042777607, + 17097689429932, + 6336398299024, + 11783538577146, + 12884677507276, + 20411681297190, + 2442525940278, + 20392672815702, + 22436836958064, + 3068913253621, + 22312349137296, + 21020455708730, + 2588466467039, + 5564028291614, + 12087393358586, + 23902761873848, + 22746289761398, + 21874695335279, + 13429537708061, + 5626946354284, + 14008870308149, + 14333634023389, + 6990951706422, + 11120693860637, + 1, + 9631788567513, + 9922523670804, + 18409285991471, + 13056718112076, + 18113807861024, + 19255994137999, + 4225156370639, + 17709907400333, + 13527720752092, + 25563730665656, + 1404110849896, + 35660600058, + 8223636593250, + 20405872151899, + 6683585055357, + 269432151142, + 18046489717861, + 6549161221813, + 8158130002826, + 5384058014184, + 1346475351431, + 13117912012789, + 19752767347849, + 19157882950615, + 10913930574837, + 22918511356826, + 9954935434006, + 10101408796117, + 21867830291000, + 3360350482866, + 12023084241499, + 13000048882520, + 3636970386063, + 5276919680157, + 6087978188082, + 14334748832270, + 23510214643892, + 24717708000254, + 12927755015193, + 1034947438382, + 24914469619014, + 10335388618567, + 5982150894261, + 5196773999112, + 7807970377766, + 3762972146817, + 25131478659917, + 10578083799641, + 6930818366526, + 2686425495504, + 21188360557486, + 18345459669164, + 1078354823587, + 24504376871621, + 34151534785, + 10526337752865, + 7079833482224, + 22454595803713, + 3721794496350, + 3605535811398, + 6679741360778, + 3309158533678, + 12998155370205, + 11625674923842, + 5420706170276, + 1020417693597, + 3227440352694, + 23256795657449, + 18348259047706, + 21864066940340, + 5541900167891, + 22839648791289, + 12540569815835, + 15120826385322, + 2707631384526, + 21630956700318, + 24495213156166, + 11693060293749, + 10455766176358, + 1177042777607, + 21426125750976, + 17097689429932, + 18417181323307, + 6336398299024, + 19624305393343, + 11783538577146, + 22376904640541, + 12884677507276, + 9910675196443, + 20411681297190, + 4768691457221, + 2442525940278, + 7880859750884, + 20392672815702, + 17577634352500, + 22436836958064, + 15327768203102, + 3068913253621, + 19806309744197, + 22312349137296, + 2332329114779, + 21020455708730, + 7051814453341, + 2588466467039, + 22215300117650, + 5564028291614, + 2517900667220, + 12087393358586, + 8573450687848, + 23902761873848, + 21441951416287, + 22746289761398, + 17301718496261, + 21874695335279, + 13715662940634, + 13429537708061, + 20684885610797, + 5626946354284, + 6890065622912, + 14008870308149, + 9214712661436, + 14333634023389, + 10394325186548, + 6990951706422, + 391989642519, + 11120693860637, + 15981907322676}; +const uint64_t con_ICTTwiddle_shoup[256] = {718658, + 718658, + 718658, + 8354886822313816593, + 718658, + 9342591962145417265, + 8354886822313816593, + 16124417262940508993, + 718658, + 193629629576812155, + 9342591962145417265, + 7602026851938233130, + 8354886822313816593, + 845891464912770977, + 16124417262940508993, + 16346808063099651946, + 718658, + 12727370490980627342, + 193629629576812155, + 13767970019466814839, + 9342591962145417265, + 743773481657582316, + 7602026851938233130, + 7564839130180332684, + 8354886822313816593, + 16413901297143266973, + 845891464912770977, + 9259679375219258149, + 16124417262940508993, + 1860222698220635140, + 16346808063099651946, + 10067589770037858187, + 718658, + 9383317769666244839, + 12727370490980627342, + 25627783285886474, + 193629629576812155, + 3869297537386599856, + 13767970019466814839, + 7259460443402106649, + 9342591962145417265, + 10301785049425030305, + 743773481657582316, + 3734704340879263837, + 7602026851938233130, + 13184115351862472500, + 7564839130180332684, + 2591147940688721251, + 8354886822313816593, + 16713687320568211617, + 16413901297143266973, + 15545264793091488509, + 845891464912770977, + 4553704709296306945, + 9259679375219258149, + 1755341339343445615, + 16124417262940508993, + 16034953067520755589, + 1860222698220635140, + 8686704569785045874, + 16346808063099651946, + 9651247636104145464, + 10067589770037858187, + 5024104894558228901, + 718658, + 7130903178041067082, + 9383317769666244839, + 13838478430545390728, + 12727370490980627342, + 18371585122310068449, + 25627783285886474, + 14664847714783607053, + 193629629576812155, + 4706608532216837363, + 3869297537386599856, + 9427295269297213551, + 13767970019466814839, + 16470576527948808437, + 7259460443402106649, + 2414943489437841691, + 9342591962145417265, + 3792301693189106998, + 10301785049425030305, + 17763583981313942892, + 743773481657582316, + 7427611965621755412, + 3734704340879263837, + 2704290856929053042, + 7602026851938233130, + 1930621759041649338, + 13184115351862472500, + 17610271812594179647, + 7564839130180332684, + 16137179803306062366, + 2591147940688721251, + 2378153974465059847, + 8354886822313816593, + 733331561164276321, + 16713687320568211617, + 15712791382753139008, + 16413901297143266973, + 10866706142809638553, + 15545264793091488509, + 8403313872161747451, + 845891464912770977, + 12287395015423751945, + 4553704709296306945, + 8468336834063846814, + 9259679375219258149, + 14669022504785789801, + 1755341339343445615, + 14655361903355193177, + 16124417262940508993, + 2205499729647419267, + 16034953067520755589, + 15106523238468485654, + 1860222698220635140, + 3998634656234161394, + 8686704569785045874, + 17177916250450275855, + 16346808063099651946, + 15720429566118828031, + 9651247636104145464, + 4043851239024468398, + 10067589770037858187, + 10300983882861344278, + 5024104894558228901, + 7991978031372397474, + 718658, + 6921964006837467151, + 7130903178041067082, + 13229984662904869926, + 9383317769666244839, + 13017636876258026113, + 13838478430545390728, + 3036443347548981496, + 12727370490980627342, + 9721807687552596546, + 18371585122310068449, + 1009075801079347101, + 25627783285886474, + 5909984018522337469, + 14664847714783607053, + 4803213324870932664, + 193629629576812155, + 12969258139462131553, + 4706608532216837363, + 5862907168989431659, + 3869297537386599856, + 967655576466393602, + 9427295269297213551, + 14195488580222160486, + 13767970019466814839, + 7843385896878253671, + 16470576527948808437, + 7154196158021867547, + 7259460443402106649, + 15715495945631677597, + 2414943489437841691, + 8640488294306456312, + 9342591962145417265, + 2613738656096997498, + 3792301693189106998, + 4375175555083456896, + 10301785049425030305, + 16895808957739968447, + 17763583981313942892, + 9290637380289205671, + 743773481657582316, + 17904988335597382938, + 7427611965621755412, + 4299121900703964996, + 3734704340879263837, + 5611262076873496215, + 2704290856929053042, + 18060943666194589619, + 7602026851938233130, + 4980889575673949034, + 1930621759041649338, + 15227189437847342816, + 13184115351862472500, + 774968555751439978, + 17610271812594179647, + 24543281126141546, + 7564839130180332684, + 5087980513157288329, + 16137179803306062366, + 2674698200028331128, + 2591147940688721251, + 4800451022174748777, + 2378153974465059847, + 9341231173959602907, + 8354886822313816593, + 3895635035931928349, + 733331561164276321, + 2319426532151452967, + 16713687320568211617, + 13186127148256612045, + 15712791382753139008, + 3982732098274565116, + 16413901297143266973, + 9012383554932625770, + 10866706142809638553, + 1945861545454566324, + 15545264793091488509, + 17603686233176214022, + 8403313872161747451, + 7514122286775477570, + 845891464912770977, + 15398061348072577806, + 12287395015423751945, + 13235658708011628465, + 4553704709296306945, + 14103168340932675686, + 8468336834063846814, + 16081346410436561470, + 9259679375219258149, + 7122388174572302675, + 14669022504785789801, + 3427059304222265515, + 1755341339343445615, + 5663644623860199625, + 14655361903355193177, + 12632311378152621334, + 16124417262940508993, + 11015426580778362105, + 2205499729647419267, + 14233967263362147469, + 16034953067520755589, + 1676147485115312125, + 15106523238468485654, + 5067844407793760411, + 1860222698220635140, + 15965207992015850665, + 3998634656234161394, + 1809510006280083302, + 8686704569785045874, + 6161380792332231910, + 17177916250450275855, + 15409434592501601182, + 16346808063099651946, + 12434022180625599182, + 15720429566118828031, + 9856873885832932621, + 9651247636104145464, + 14865363029917218071, + 4043851239024468398, + 4951602281574965200, + 10067589770037858187, + 6622228979459311253, + 10300983882861344278, + 7469967214534403526, + 5024104894558228901, + 281706577916525827, + 7991978031372397474, + 11485529034691002640}; + +const uint64_t con_sample[256] = { + 34359214082, 68718428164, 103077642246, 137436856328, 171796070410, + 206155284492, 240514498574, 274873712656, 309232926738, 343592140820, + 377951354902, 412310568984, 446669783066, 481028997148, 515388211230, + 549747425312, 584106639394, 618465853476, 652825067558, 687184281640, + 721543495722, 755902709804, 790261923886, 824621137968, 858980352050, + 893339566132, 927698780214, 962057994296, 996417208378, 1030776422460, + 1065135636542, 1099494850624, 1133854064706, 1168213278788, 1202572492870, + 1236931706952, 1271290921034, 1305650135116, 1340009349198, 1374368563280, + 1408727777362, 1443086991444, 1477446205526, 1511805419608, 1546164633690, + 1580523847772, 1614883061854, 1649242275936, 1683601490018, 1717960704100, + 1752319918182, 1786679132264, 1821038346346, 1855397560428, 1889756774510, + 1924115988592, 1958475202674, 1992834416756, 2027193630838, 2061552844920, + 2095912059002, 2130271273084, 2164630487166, 2198989701248, 2233348915330, + 2267708129412, 2302067343494, 2336426557576, 2370785771658, 2405144985740, + 2439504199822, 2473863413904, 2508222627986, 2542581842068, 2576941056150, + 2611300270232, 2645659484314, 2680018698396, 2714377912478, 2748737126560, + 2783096340642, 2817455554724, 2851814768806, 2886173982888, 2920533196970, + 2954892411052, 2989251625134, 3023610839216, 3057970053298, 3092329267380, + 3126688481462, 3161047695544, 3195406909626, 3229766123708, 3264125337790, + 3298484551872, 3332843765954, 3367202980036, 3401562194118, 3435921408200, + 3470280622282, 3504639836364, 3538999050446, 3573358264528, 3607717478610, + 3642076692692, 3676435906774, 3710795120856, 3745154334938, 3779513549020, + 3813872763102, 3848231977184, 3882591191266, 3916950405348, 3951309619430, + 3985668833512, 4020028047594, 4054387261676, 4088746475758, 4123105689840, + 4157464903922, 4191824118004, 4226183332086, 4260542546168, 4294901760250, + 4329260974332, 4363620188414, 4397979402496, 4432338616578, 4466697830660, + 4501057044742, 4535416258824, 4569775472906, 4604134686988, 4638493901070, + 4672853115152, 4707212329234, 4741571543316, 4775930757398, 4810289971480, + 4844649185562, 4879008399644, 4913367613726, 4947726827808, 4982086041890, + 5016445255972, 5050804470054, 5085163684136, 5119522898218, 5153882112300, + 5188241326382, 5222600540464, 5256959754546, 5291318968628, 5325678182710, + 5360037396792, 5394396610874, 5428755824956, 5463115039038, 5497474253120, + 5531833467202, 5566192681284, 5600551895366, 5634911109448, 5669270323530, + 5703629537612, 5737988751694, 5772347965776, 5806707179858, 5841066393940, + 5875425608022, 5909784822104, 5944144036186, 5978503250268, 6012862464350, + 6047221678432, 6081580892514, 6115940106596, 6150299320678, 6184658534760, + 6219017748842, 6253376962924, 6287736177006, 6322095391088, 6356454605170, + 6390813819252, 6425173033334, 6459532247416, 6493891461498, 6528250675580, + 6562609889662, 6596969103744, 6631328317826, 6665687531908, 6700046745990, + 6734405960072, 6768765174154, 6803124388236, 6837483602318, 6871842816400, + 6906202030482, 6940561244564, 6974920458646, 7009279672728, 7043638886810, + 7077998100892, 7112357314974, 7146716529056, 7181075743138, 7215434957220, + 7249794171302, 7284153385384, 7318512599466, 7352871813548, 7387231027630, + 7421590241712, 7455949455794, 7490308669876, 7524667883958, 7559027098040, + 7593386312122, 7627745526204, 7662104740286, 7696463954368, 7730823168450, + 7765182382532, 7799541596614, 7833900810696, 7868260024778, 7902619238860, + 7936978452942, 7971337667024, 8005696881106, 8040056095188, 8074415309270, + 8108774523352, 8143133737434, 8177492951516, 8211852165598, 8246211379680, + 8280570593762, 8314929807844, 8349289021926, 8383648236008, 8418007450090, + 8452366664172, 8486725878254, 8521085092336, 8555444306418, 8589803520500, + 8624162734582, 8658521948664, 8692881162746, 8727240376828, 8761599590910, + 8795958804992}; + +const uint64_t con_NegModn_shoup[256] = { + 12191655058664, 304724738883, 4600774299129309644, + 13845970226104576361, 4970592428225372796, 8367062477718566923, + 3294236965032321920, 1814852574693535748, 7511474276989238501, + 18014075139885811565, 9357451159425740401, 1714729418009381772, + 2407310681514169103, 3255906620457763194, 2909719840234738179, + 10169566107912123188, 12514880144198686037, 1543992702915071157, + 9682986526057329499, 15244152151508641447, 10509625003593325167, + 1410389955457458458, 13296329095340384970, 12325278380567008254, + 17905923856117641145, 6620982084000050408, 9351083717174626410, + 11734110725029888999, 4224816751578956142, 2625943516451546951, + 1752353384843690777, 16831105739624403307, 4399860465521324352, + 16809565765870314894, 11796660401458932194, 8537846329819547834, + 5376734085819115454, 6484042594515760956, 13081601959926560416, + 5434995939711627003, 6987757206824195827, 16118330435929098024, + 11093833146808109348, 5480081909965564390, 12016862044513080442, + 18099915168551346567, 2750819843782152308, 10232597519030357907, + 554724057819924566, 17979588076442207789, 8570125218137571368, + 10997424393628469834, 9351311574813358130, 1862400941240064962, + 3297934208333983027, 6919310200239057422, 1636736211265079070, + 8635680642414299283, 10325294813721584324, 6757539580437106919, + 4409071270848205561, 1471236991845051217, 17884002611564235460, + 11347273727820659807, 16418996611505883996, 3296333144817683650, + 14248404754136794887, 4681427745090460643, 792395075454373547, + 16558724889143627759, 5329081931732904638, 300140064237077778, + 16481313085265638842, 18382528052312011157, 11058730690995786859, + 13504792747209916835, 9466927813324202819, 10955282402857414204, + 13348649940185202876, 7428120034178046173, 17231882062385015464, + 4997089139069035495, 15901700215762474078, 326712011129292668, + 8114017759934469399, 6488711722730249162, 9738511466090230506, + 16930847249893179884, 5305472940202757805, 716090222718723514, + 4093019076946933125, 5650881776324313453, 680988178254147745, + 13188789200767441295, 4383884034249417748, 441615763999978231, + 4780505812788444055, 13616774095716356046, 13050740308232766441, + 2837651977200424275, 10873725157195869345, 9938556701178823284, + 17850253625778571486, 10591737275702331494, 10631750208849271364, + 3002318639721944456, 11887505735777101904, 8729342347082469412, + 8147182366476910909, 6596721716225925143, 946034149692367208, + 9390114561343839047, 3024780877789220988, 17968348457900709359, + 8926405456749076021, 14342131137248300085, 6226396796452806705, + 11429075857585088688, 765566210003514546, 6159509894324823282, + 15769932009058851850, 6439177054156801093, 2132713605722678087, + 14477720401172538612, 15040243750037046321, 6198458793881693269, + 3894253511161480489, 1294633678438752483, 12433583837313694208, + 3732017642113492528, 11091538299144254929, 17140782162622302950, + 17637786178978938028, 9218295175332253448, 16901099987116510395, + 14208143658117369351, 6055601559876429644, 3551409970739999461, + 11745204925207510934, 10065496662198730923, 386723265195494282, + 15252480886229162013, 13673785214002294189, 349072811423478924, + 7750383240885128874, 10758253463874854488, 492486592317679846, + 7536567183152110378, 4360476833294204099, 12824205994597391029, + 14748915621378087848, 14123016985885844121, 2793609008017093047, + 6930623609572936311, 9231991897624427333, 16825470192920516160, + 14477293411762871884, 8950302490020838671, 13115738448022820516, + 15721920356660658831, 14501812712237346308, 13131073155079057421, + 7041935637947255676, 18198158343365220965, 8633361406577886033, + 3141747837174452181, 12763015537032429017, 16418703019073133367, + 14385115189471756547, 872489449757377314, 17061992011429906895, + 3135342477512516453, 3654200776070469906, 4265990795126539828, + 7716123640980307327, 17745281843237425989, 5074688993219525368, + 13275634691656038198, 11269101204359653431, 2243299061945840246, + 16684344991410180403, 12831285018168688980, 17878546673140284269, + 4417217338765653899, 11833699664361673218, 6515253005036860915, + 16122288707027505609, 13857483328334852139, 4764240769650421548, + 1626915116124753068, 6520222408545797705, 8604261716116819280, + 8693309731754780212, 12303014204411364368, 14893164877069048615, + 1811931660652517971, 11451253903841874940, 1703008489023478915, + 15485987741499170042, 15653981949075597773, 8930860977932092080, + 7620886358657019844, 7757995594551923483, 586013812628744162, + 10372431387153085887, 17661742493721737763, 7795118898959331362, + 1406150770257422709, 1709547658982075612, 15696592109703877807, + 17722061145880775200, 5958917464235183829, 2601080446123983761, + 15004692087751375295, 15509663769362237107, 11054359485782353616, + 15930397849436149527, 16806803705895194828, 10099836596531255677, + 17094570995929754057, 16266135194879135780, 9491571856247046989, + 8368156193856409121, 3671014988807075698, 9163240728618840082, + 18373263813455202435, 14460227365580342295, 4758821249478879388, + 15244663175901151703, 4823441675767427361, 1925883441203775731, + 12540040873137976443, 13244433890703482144, 8569879106902783268, + 1721790325798116494, 12886366124261129418, 16932041145357395096, + 6118942568472703450, 10426103355458156024, 8075594936512397486, + 1310149414082877043, 3679410473945937808, 12812060676337437345, + 9772544045039058666, 4594674563340394751, 16298362699047623905, + 12165280149661691064, 836048739423488211, 3281553100841658751, + 15597837648211678009, 2185444140563025981, 15557848222006373029, + 1761082975822988479, 12019392191068039298, 16998434045378741685, + 17309128569570633760}; +const uint64_t con_modn_shoup[256] = { + 6707306413801851557, 13460171823992435485, 3431527379264082046, + 2655838758876326921, 1394206572926674797, 6280840331630052466, + 9819686826703062089, 897387507355885984, 9415375063653129697, + 12141032226860303707, 15645796011122766097, 9511787522468104474, + 662323987270685347, 2148154270741155019, 160825845275050947, + 16665213678424199989, 16285944302335888767, 13397928502642403480, + 3412774470543106878, 8968505167659171061, 6468346824267116460, + 18207446362621740375, 4888461093714324027, 5256801328785751106, + 4818132091056833704, 10673209946648588774, 6637502763148359109, + 8048600301697596756, 1854925695892386622, 15851894579959613623, + 13445981410741012143, 10270107099676478705, 10916204078951271193, + 17379104895483187833, 15999685066846132052, 567702515627661427, + 18226486475591665215, 15934882438721348206, 14257961411892698709, + 3825405099016186727, 3589153149955596935, 4933298482761233834, + 16539253555492019161, 16199700734304713500, 1327731940547971548, + 10610469855838191242, 4608147056152067027, 14479365299895066733, + 9788904995561131291, 6706337193109549621, 6677843030700129921, + 3120281824041023667, 3275432682716744590, 1344896817422862493, + 14926451346553767012, 5132828290524441432, 8121031984155373707, + 13784599574194620378, 4739244628103651205, 7124941662936329775, + 8968236967802172484, 2988744753600565844, 6800797399003367084, + 1379445650297039370, 1430731647792757132, 13610441310698150243, + 11520787325928743026, 9924325084065238870, 17631138561608068144, + 18078381803587572251, 17643556967221822627, 17796856855673094311, + 1204426651706444893, 6095086536941196754, 3969923057433735496, + 4630457334495416233, 10789086249650894227, 7507586610082244831, + 18267030682341867913, 16066828480903004065, 14170393815716935927, + 2170632152020087183, 13662925123207067636, 6936670780491359428, + 10321519377799964174, 2279004717002885385, 1812947213826551145, + 14382495392705155412, 17568420724181620094, 9130554643424788581, + 3844701922568660716, 13144352407010793993, 15088128657245239855, + 11228458894019541307, 13834262665308504417, 6557198421479987348, + 1958003609747250778, 17381102770990458164, 1766258802946908157, + 540545484347655331, 17577555667475583068, 12134637792304556632, + 11135400007937208986, 5672425062848442456, 15531450524884383930, + 14433405663987689154, 18246542000065810082, 13286664354856757936, + 8294196081264154840, 5982513056487063902, 946294328100498060, + 9725521233313374942, 13779754927669634447, 1996718839028364426, + 3936684418657061334, 2346553974171991350, 13290910364739941671, + 8221193909388897806, 13596916242869969471, 10434471876112363323, + 14246769921345152184, 1622660458487813999, 8878779360681631410, + 18348852439447001620, 16028470176130612477, 12841757037633397004, + 11086377804137065556, 12758752772414795765, 6876076611198679861, + 16988111019667941961, 1765899717999946292, 14733583555917578092, + 16191473330361291759, 13840441490042280846, 6486902021132720932, + 3996411524551136973, 964303974127137183, 3876461103860342055, + 5304779994154416444, 5107418024230673581, 16191604327022634802, + 16803848964540714708, 10112804462803046317, 6463768302535329612, + 1680373001655953385, 7137906115406951582, 970838729085150161, + 9556435610243251174, 17446289842214945949, 9665162289252207113, + 7604198540368551646, 1220486027779675578, 5364440337208143244, + 14241738412098749367, 2286879042508078657, 9653688497919661422, + 16073537475764784162, 12477583642553505095, 8359381430158825600, + 12362816776698324191, 4348426194782617046, 2663420380941202259, + 13461314021014160892, 11871859441537656244, 8423310865757369673, + 5889990515471957850, 3931031807819832738, 17976439351416312479, + 1633093626155221259, 3455171914490611494, 14988223176342893715, + 17130123970279794665, 4702459268974675817, 1091430039346991150, + 5156823928441090932, 10089406231387272978, 16925187590423810791, + 7577884842347519915, 10381738316792437595, 16528541544495675013, + 5815404723819793441, 6097101853643090337, 458364875272336803, + 7005810603059918949, 14052006965745796071, 12731065718772976810, + 3983003481666015622, 1472481138684565791, 7086540843869480347, + 762122247871102630, 15006093430418347741, 4149542905368836311, + 11483004066878431355, 3281851005022492305, 5353581131102023978, + 10142620438696769848, 8386367143613751605, 6150798453949235440, + 8044933108941272325, 2824428019982542492, 8253034621114380355, + 17267400824644797298, 7819465119641314663, 11538526436656806501, + 17291333536799021270, 7794620899865412501, 1490921524009742645, + 12114551713343966263, 1861754564954994608, 18421658786333503965, + 5619178818445504794, 1033900325738606568, 15388445472761661481, + 1582355764125440406, 6415798796592853541, 2814038840809724531, + 3813460348665722975, 11144908283381828079, 17406626293698576236, + 14758333830406625705, 10554974721914125925, 3829052279593923953, + 8171294983270628064, 6423782500177742751, 15551589171272452013, + 8021071333480059931, 3250514172507803918, 17393638039278558421, + 4517309645980754291, 2203240740930873265, 3691584123813477930, + 16359699686913203430, 3171815310279072404, 2816445194756100328, + 17972521234261373176, 8450093692810782746, 11930625254294164192, + 11683058118729788538, 17275944271238937102, 2758643121997351267, + 12564301311367916072, 10288670342937014961, 12603667279289376042, + 2289082877377123131, 8898027755861140854, 9678253529564474977, + 7313437916645855855, 13699047632633126564, 6150848657898123031, + 11838689987468026615, 3634191852514867107, 1905424774787734347, + 3648272932716999785, 1119422476116358561, 5756877321311841064, + 852645000108100762}; + +const uint64_t con_modn[256] = { + 9333096382970, 18729587290981, 4774908703376, 3695551922783, + 1940013400330, 8739676490778, 13663917815893, 1248698595578, + 13101325260773, 16894028238946, 21770844084239, 13235481465235, + 921611930123, 2989118077722, 223786274582, 23189345455199, + 22661598931668, 18642976827478, 4748814367773, 12479513828187, + 9000588406866, 25335334530182, 6802205794289, 7314744614429, + 6704344250500, 14851579904381, 9235965857057, 11199482732617, + 2581095783456, 22057626540366, 18709841585862, 14290669534266, + 15189701874269, 24182712259205, 22263274347921, 789947852137, + 25361828505338, 22173102655017, 19839697170840, 5322982474687, + 4994242131662, 6864596215677, 23014074206920, 22541592556055, + 1847515143543, 14764278195098, 6412153846657, 20147776703637, + 13621085450832, 9331747729788, 9292098614672, 4341816104581, + 4557705737257, 1871399761360, 20769888905990, 7142238358773, + 11300270117575, 19181022677281, 6594573762956, 9914228329630, + 12479140632904, 4158784634030, 9463187409381, 1919473841660, + 1990837385860, 18938684582251, 16030968601544, 13809520071263, + 24533412578018, 25155743514836, 24550692566342, 24764006607202, + 1675938049209, 8481203422399, 5524076617666, 6443198198184, + 15012819698376, 10446671899648, 25418245040684, 22356703202399, + 19717848433849, 3020395645723, 19011714864524, 9652252778856, + 14362208795691, 3171194123063, 2522683479774, 20012983967979, + 24446141833919, 12705002762615, 5349833658007, 18290130245715, + 20994859979403, 15624198840008, 19250127976405, 9124223808234, + 2724526848866, 24185492267083, 2457717394758, 752159330967, + 24458852943410, 16885130498133, 15494709071759, 7893077574002, + 21611734394271, 20083824695903, 25389735407361, 18488154764613, + 11541224847936, 8324559446647, 1316751558069, 13532888084307, + 19174281438672, 2778398394920, 5477825648916, 3265187700941, + 18494063011089, 11439643637202, 18919864693336, 14519380167984, + 19824124412790, 2257902878082, 12354661979636, 25532098545828, + 22303328310387, 17869074225130, 15426495555819, 17753575281189, + 9567937080815, 23638651304926, 2457217734505, 20501516840048, + 22530144271267, 19258725700024, 9026407629974, 5560934843814, + 1341811656985, 5394025988244, 7381506065388, 7106880429702, + 22530326550768, 23382253965349, 14071785740883, 8994217476126, + 2338208845137, 9932268132346, 1350904651117, 13297608477304, + 24276198905654, 13448899697950, 10581105665079, 1698284382520, + 7464522360993, 19817123158061, 3182151105547, 13432934123429, + 22366038648227, 17362333487984, 11631929089916, 17202637455849, + 6050756933756, 3706101613688, 18731176637618, 16519479143661, + 11720885751106, 8195815992924, 5469960142481, 25013892423961, + 2272420443467, 4807809649387, 20855843297456, 23836259774544, + 6543387596550, 1518705293012, 7175627857870, 14039227522200, + 23551094483506, 10544490626989, 14446002367510, 22999168636052, + 8092031203719, 8484007699386, 637806489753, 9748459599914, + 19553115258847, 17715049243424, 5542277793005, 2048931052437, + 9860794279741, 1060479416946, 20880709434976, 5774014409329, + 15978393875324, 4566636717445, 7449411970774, 14113274135717, + 11669479225410, 8558725554092, 11194379899775, 3930140851030, + 11483949415573, 24027278051315, 10880645243031, 16055652255452, + 24060579990344, 10846074957521, 2074590517349, 16857181080091, + 2590598032081, 25633407269608, 7818986381445, 1438653730005, + 21412745433223, 2201819619921, 8927468770344, 3915684494845, + 5306361569153, 15507939666231, 24221007790420, 20535956402379, + 14687057713172, 5328057460942, 11370210177906, 8938577950461, + 21639756958944, 11161176667880, 4523032077944, 24202934867719, + 6285755222231, 3065769911403, 5136772991688, 22764228224346, + 4413524024920, 3919032893133, 25008440434326, 11758153573817, + 16601250716427, 16256765494651, 24039166103774, 3838602347261, + 17482999566320, 14316499953788, 17537776603502, 3185217702177, + 12381445775745, 13467115923670, 10176517480470, 19061978687476, + 8558795411985, 16473324444188, 5056912677186, 2651364344560, + 5076506247331, 1557656265913, 8010591374531, 1186440200509}; + +const uint64_t con_negmodn[256] = { + 16964469, 424019, 6401894787706, 19266418836935, + 6916490070832, 11642617129611, 4583871558294, 2525334755192, + 10452081518310, 25066261947517, 13020725188601, 2386015181305, + 3349729567758, 4530535572395, 4048822886706, 14150768550371, + 17414230900019, 2148438109411, 13473701803265, 21211963868823, + 14623954394803, 1962532286399, 18501603077376, 17150403451982, + 24915771379012, 9212977628804, 13011865007095, 16327804279195, + 5878756614091, 3653953229995, 2438368255268, 23420181278230, + 6122326797510, 23390208791701, 16414841030287, 11880259791803, + 7481628891337, 9022428789204, 18202814125863, 7562699214397, + 9723338623368, 22428367247623, 15436869531216, 7625435532069, + 16721247660836, 25185706805352, 3827716396164, 14238475626629, + 771888561246, 25018273818270, 11925175285031, 15302718483079, + 13012182066906, 2591497453070, 4589016205930, 9628095843305, + 2277489035211, 12016394480257, 14367462246953, 9402994931941, + 6135143467657, 2047199844460, 24885268361015, 15789527546821, + 22846738829691, 4586788354809, 19826399253251, 6514122608881, + 1102603512699, 23041168217908, 7415321774962, 417639507590, + 22933451083458, 25578957556172, 15388025094543, 18791676504056, + 13173059985833, 15244078659971, 18574406593072, 10336095586836, + 23977854331699, 6953359767977, 22126930187864, 454613894345, + 11290509950577, 9028925797237, 13550963759856, 23558969797986, + 7382470287366, 996426680878, 5695362517489, 7863100477813, + 947582816588, 18351963245953, 6100095880400, 614500402303, + 6651987960218, 18947496538840, 18159870691938, 3948541749940, + 15130593217749, 13829323110816, 24838307253570, 14738212146356, + 14793889424027, 4177672452723, 16541250681021, 12146723059739, + 11336657904803, 9179219762866, 1316389524488, 13066175725596, + 4208928253476, 25002634090981, 12420932836746, 19956817831625, + 8663919289617, 15903353740873, 1065271628385, 8570847366880, + 21943577095355, 8959999195771, 2967632669729, 20145487856614, + 20928228991240, 8625044060925, 5418783803376, 1801459506404, + 17301108549417, 5193035506044, 15433676291917, 23851090457722, + 24542662618170, 12827092136570, 23517577016220, 19770376661401, + 8426260786132, 4941723175812, 16343241659443, 14005957786192, + 538118476320, 21223553153567, 19026826485772, 485728546259, + 10784519049747, 14969916430857, 685286246050, 10486997846460, + 6067525180764, 17844652264037, 20522851133707, 19651924426229, + 3887256749521, 9643838249156, 12846150879588, 23412339516230, + 20144893708521, 12454185129251, 18250314435812, 21876769743537, + 20179011874660, 18271652405204, 9798726937000, 25322410418239, + 12013167305254, 4371685676235, 17759506841577, 22846330301095, + 20016629368117, 1214053395705, 23741455377745, 4362772733333, + 5084754734868, 5936049556026, 10736847431928, 24692240903905, + 7061338571003, 18472807186409, 15680751884737, 3121510345477, + 23215966333539, 17854502598167, 24877676520710, 6146478574826, + 16466380499233, 9065857514772, 22433875110749, 19282439111864, + 6629355444881, 2263823157752, 9072772350483, 11972675608732, + 12096584322810, 17119423250821, 20723571381669, 2521270355776, + 15934214093827, 2369705719179, 21548473747896, 21782234669846, + 12427132614432, 10604326464443, 10795111503124, 815427692965, + 14433051942109, 24575997399777, 10846767914314, 1956633536347, + 2378804856705, 21841525942900, 24659929720676, 8291726603906, + 3619356714281, 20878759491829, 21581418515345, 15381942641014, + 22166862429589, 23386365434864, 14053741187924, 23786788449419, + 22634034914498, 13207351728848, 11644139015969, 5108151409836, + 12750484862981, 25566066534228, 20121146611837, 6621814279920, + 21212674950085, 6711732442242, 2679832211353, 17449241602344, + 18429391832383, 11924832825669, 2395840307649, 17931146967803, + 23560631082077, 8514398661719, 14507735521828, 11237045282023, + 1823049373824, 5119833576588, 17827752271860, 13598319482093, + 6393407114933, 22678878907605, 16927768787928, 1163346801906, + 4566222189129, 21704110883785, 3041006261694, 21648466315601, + 2450515324295, 16724768314315, 23653015609757, 24085341458771}; + +// const uint64_t +// con_NegModn_shoup[128]={6045753221356,50334102921,15614544885764143922,2832199278545211903,12996214100276393278,3548646091445667944,14036887970724850822,6311739757749710492,8931462630315306839,16125031722383780835,4173271538509672110,6151594547312679773,12393126291935972955,7185341161223532025,17608762589683592848,1218385367359393750,1485890938396535942,1711368403314332280,16342142035425371378,2974457136671677393,3818636342861149169,3468246016662833290,1662109691301686949,12939473651793102839,2218899434480232327,16300960611552540028,15363593521069914409,4415773996530236314,11907537573327410481,7128093048255283088,13479018750867177886,13911008231998953321,3070793311218315214,14318967785635216799,8664781961319673479,12649834620626344384,13559706796003714285,15077852801214906302,3657533299403085759,1863145962152572116,17842520749508039838,12614919159661470796,2520847433321131727,7259210960235177489,17811155232180342752,17844022889617881171,8723343436501783654,1512398760398376002,7738850058338842939,9091791665834107325,14580422593189409977,3579321992548823770,8077545311409309505,13661970680565413799,966451245035613908,4960458923731683758,10572427442137126404,18389607259989890020,1443313362982319469,4114120977522679446,5676292866965769987,10625311199154638451,7789827231482564476,14889160180296186732,3583090055705480708,7589829412062521948,9965368077161349221,12090061775956438515,2426287668361913247,5174571361352051931,7974826356436538611,5989152965092226049,8960675188926997709,16393110776837529949,502625137986763073,16854264088531148148,10544404935119688617,7587554654844883926,826160461184423407,10482740524943004940,18189233974710725326,3462743284132793878,17554832001094711340,16560733667977430792,238674138476223640,6678166015841481650,15505834580252951451,8068709492937881797,4458157699082948211,7131083097086064243,16285751472251770271,13967538606001587051,13843638710010451422,12052128430605594902,17939900570363940835,9709291751168401039,1970332567150533519,6402150135342276010,6839993727420372042,3014158244276600070,1439452232107449647,841856549580904979,14214132554715378971,6629910273084033084,7874327013830067781,8678974265337152318,15531311665009776889,10951342171909979000,3956834127023368956,3214250517008740625,5658610877251100021,1974447772720527565,10065913589073172314,3794371718823667350,14658505701237349701,4586216326212101511,9883687596712292699,10250926880044021830,4571774906588313675,14486854841753579654,445069001901439648,7032901876659516529,7085835865947517875,8597216160938791262,9268845303974687933,12870449106922623958,15723888207682576782,12296640027273446714}; +// const uint64_t +// con_modn_shoup[128]={12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,12656799714631949916,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,5789944359081913648,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,1949458368552867584,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,16497285705160995980,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,15477458790088929150,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,2969285283624934415,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,1699167585985028950,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,16747576487728834615,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,8587268584366811608,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,9859475489347051957,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,13799404763057967906,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,4647339310655895658,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9378027243525808005,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,9068716830188055559,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,2620594594824921088,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477,15826149478888942477}; + +void genTwiddle_shoup(uint64_t *twiddle, uint64_t *twiddle_shoup, int length) { + for (int i = 0; i < length; i++) { + __uint128_t temp = twiddle[i]; + temp <<= 64; + twiddle_shoup[i] = temp / MOD; + } +} + +__forceinline__ __device__ void csub_q(uint64_t &operand, + const uint64_t modulus) { + const uint64_t tmp = operand - modulus; + operand = tmp + (tmp >> 63) * modulus; +} + +[[nodiscard]] __inline__ __device__ uint64_t multiply_and_reduce_shoup_lazy( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t &operand2_shoup, const uint64_t &modulus) { + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__forceinline__ __device__ uint64_t precompute_shoup(uint64_t b, uint64_t mod) { + // 计算 b_shoup = floor((b << 64) / mod) + uint64_t q = 0; + uint64_t remainder = + b; // 初始余数 = b (因为 b << 64 等价于高位为 b,低位为 0) + + // 通过逐位除法模拟 128 位运算 + for (int i = 0; i < 64; i++) { + remainder <<= 1; // 左移一位(模拟 128 位左移) + q <<= 1; // 商左移一位 + + if (remainder >= mod) { + remainder -= mod; + q |= 1; // 当前位商为 1 + } + } + return q; +} + +[[nodiscard]] __inline__ __device__ uint64_t +multiply_and_reduce_shoup_lazy_dotproduct( + const uint64_t &operand1, const uint64_t &operand2, + const uint64_t & + modulus) { /* + uint64_t hi_1=operand2; + uint64_t lo=0; + uint64_t q_hi=hi_1/modulus; + uint64_t r_hi=hi_1%modulus; + uint64_t combined_lo=(r_hi<<32)|(lo>>32); + uint64_t q_lo=combined_lo/modulus; + uint64_t operand2_shoup=(q_hi<<32)+q_lo; + const uint64_t hi=__umul64hi(operand1,operand2_shoup); + uint64_t + result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; + */ + uint64_t operand2_shoup = precompute_shoup(operand2, modulus); + const uint64_t hi = __umul64hi(operand1, operand2_shoup); + uint64_t result = operand1 * operand2 - hi * modulus >= modulus + ? operand1 * operand2 - hi * modulus - modulus + : operand1 * operand2 - hi * modulus; + return result; +} + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, uint64_t mod); + +__global__ void multiply_dotproduct(uint64_t *inout, uint64_t *inoutA, + uint64_t *result, int length, + uint64_t mod) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < length; + i += blockDim.x * gridDim.x) { + result[i] = + multiply_and_reduce_shoup_lazy_dotproduct(inout[i], inoutA[i], mod); + } +} + +/* +uint64_t multiply_and_reduce_shoup_lazy_dotproduct_host(const uint64_t +&operand1, const uint64_t &operand2, const uint64_t &modulus) +{ + __uint128_t temp=operand2; + temp<<=64; + uint64_t operand2_shoup=temp/MOD; + const uint64_t hi=__umulh(operand1,operand2_shoup); + uint64_t +result=operand1*operand2-hi*modulus>=modulus?operand1*operand2-hi*modulus-modulus:operand1*operand2-hi*modulus; + return result; +} +*/ +__device__ __forceinline__ void ct_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + const uint64_t tw_y = multiply_and_reduce_shoup_lazy(y, tw, tw_shoup, mod); + // const uint64_t hi = __umul64hi(y, tw_shoup); + // const uint64_t tw_y = y * tw - hi * mod; //得到tw_y=tw*y % mod [0,2*mod) + // tw_y>=mod?tw_y-mod:tw_y; + // csub_q(x, mod2); + // const uint64_t mod2 = 2 * mod; + // const uint64_t tmp = x - mod2; + // x = tmp + (tmp >> 63) * mod2; //得到x=x-mod2 [0,2*mod) + // y = x + mod2 - tw_y; //y:[0,4*mod-1] + y = x - tw_y + mod; + x += tw_y; // x:[0,4*mod-1] + csub_q(x, mod); + csub_q(y, mod); +} + +__device__ __forceinline__ void gs_butterfly(uint64_t &x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + uint64_t v = x - y + mod; + x = u; + csub_q(x, mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +__device__ __forceinline__ void gs_butterfly_left(uint64_t &x, uint64_t y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + uint64_t u = x + y; + // uint64_t v=x-y+mod; + x = u; + csub_q(x, mod); + // y=multiply_and_reduce_shoup_lazy(v,tw,tw_shoup,mod); +} + +__device__ __forceinline__ void gs_butterfly_right(uint64_t x, uint64_t &y, + const uint64_t &tw, + const uint64_t &tw_shoup, + const uint64_t &mod) { + /* + const uint64_t mod2 = 2 * mod; + const uint64_t t = x + mod2 - y; // [0, 4q) + uint64_t s = x + y; // [0, 4q) + csub_q(s, mod2); // [0, 2q) + x = s; + y = multiply_and_reduce_shoup_lazy(t, tw, tw_shoup, mod); + */ + // uint64_t u=x+y; + uint64_t v = x - y + mod; + // x=u; + // csub_q(x,mod); + y = multiply_and_reduce_shoup_lazy(v, tw, tw_shoup, mod); +} + +/* +128线程的CT变换和NCT变换 + + + + +*/ + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer[256]; + //__shared__ uint64_t s_twiddle[256]; + + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + register uint64_t pairsInGroup; + register uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + register uint64_t omgn, omgn_shoup; + register uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + // blfIdx=glbIdx; + + twiddleId = pairsInGroup + j; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("pairsInGroup:%llu ",pairsInGroup); + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + } + // uuint64_t start=clock64(); + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // uuint64_t end=clock64(); + // d_clocktime[tid]=end-start; + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + // printf("inout[glb]:%llu ",buffer[glbIdx]); + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup); + +__global__ void TestInvctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *ICTTwiddle, + uint64_t *ICTTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup) { + // extern __shared__ uint64_t s_twiddle[]; + + __shared__ uint64_t buffer4[256]; + // uint64_t inv=ksm(n,mod-2,mod); + //__shared__ uint64_t s_twiddle[256]; + + // uint64_t invn=inv; + // uint64_t invn_shoup=inv_shoup; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + register uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t ICTtwiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + ICTtwiddleId = pairsInGroup + j; + omgn = ICTTwiddle[ICTtwiddleId]; + omgn_shoup = ICTTwiddle_shoup[ICTtwiddleId]; + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } + + else { + samples[0] = buffer4[glbIdx]; + samples[1] = buffer4[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == 1) { + // csub_q(samples[0],mod); + // csub_q(samples[1],mod); + inout[glbIdx] = samples[0]; + inout[glbIdxadd] = samples[1]; + + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + __syncthreads(); + + } else { + buffer4[glbIdx] = samples[0]; + + buffer4[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample); + +__global__ void XYfixWarp(uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t *ICTTwiddle, uint64_t *ICTTwiddle_shoup, + uint64_t *NCTtwiddle, uint64_t *NCTtwiddle_shoup, + uint64_t *InvNCTtwiddle, + uint64_t *InvNCTtwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *sample) { + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < 32; + i += blockDim.x * gridDim.x) { + // printf("tid:%llu ",i); + // ct x fix + /* + if (i == 0) { + printf(">>> [GPU 接收数据检查] <<<\n"); + printf("inout[0] = %llu\n", (unsigned long long)inout[0]); + printf("inoutA[0] = %llu\n", (unsigned long long)inoutA[0]); + printf("modn[0] = %llu\n", (unsigned long long)modn[0]); + printf("negmodn[0] = %llu\n", (unsigned long long)negmodn[0]); + printf("twiddle[3] = %llu\n", (unsigned long long)twiddle[3]); + printf("mod = %llu\n", (unsigned long long)mod); + } + */ + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + // uint64_t numIdx[8]; + register uint64_t C; + register uint64_t omegnIdx[8]; + register uint64_t left[8]; + register uint64_t right[8]; + + // __shared__ uint64_t s_tile[8][32]; + /* + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inout[i*8+j]; + } + __syncwarp(); + #pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + /* + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],twiddle[omegnIdx[0]],twiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],twiddle[omegnIdx[1]],twiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],twiddle[omegnIdx[2]],twiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],twiddle[omegnIdx[3]],twiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],twiddle[omegnIdx[4]],twiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],twiddle[omegnIdx[5]],twiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],twiddle[omegnIdx[6]],twiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],twiddle[omegnIdx[7]],twiddle_shoup[omegnIdx[7]],mod); + + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + //printf("regs[0]:%llu ",regs[0]); + */ + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + /* + inout[8*i]=regs[0]; + inout[8*i+1]=regs[1]; + inout[8*i+2]=regs[2]; + inout[8*i+3]=regs[3]; + inout[8*i+4]=regs[4]; + inout[8*i+5]=regs[5]; + inout[8*i+6]=regs[6]; + inout[8*i+7]=regs[7]; + */ + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + /* + +#pragma unroll + for (int j = 0; j < 8; j++) { + + s_tile[j][i] = inoutA[i*8+j]; + } + __syncwarp(); +#pragma unroll + for (int j = 0; j < 8; j++) { + + regs[j] = s_tile[j][i]; + } + */ + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + /* + inoutA[8*i]=regs[0]; + inoutA[8*i+1]=regs[1]; + inoutA[8*i+2]=regs[2]; + inoutA[8*i+3]=regs[3]; + inoutA[8*i+4]=regs[4]; + inoutA[8*i+5]=regs[5]; + inoutA[8*i+6]=regs[6]; + inoutA[8*i+7]=regs[7]; + */ + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // x*negmodn fix + register uint64_t Aa1[8]; + /* + __shared__ uint64_t s_negmodtile[8][32]; + __shared__ uint64_t s_negmodshouptile[8][32]; + + #pragma unroll + for (int j = 0; j < 8; j++) { + + s_negmodtile[j][i] = negmodn[i*8+j]; + s_negmodshouptile[j][i] = negmodn_shoup[i*8+j]; + } + __syncwarp(); + */ +#pragma unroll + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("negmodn_shoup[0] = %llu\n", (unsigned long + long)negmodn_shoup[8*i+j]); printf("negmodn[0] = %llu\n", (unsigned long + long)negmodn[8*i+j]); printf("mod = %llu\n", (unsigned long long)mod); + } + } + */ + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * i + j], + negmodn_shoup[8 * i + j], mod); + // Aa1[j]=multiply_and_reduce_shoup_lazy(CTinout[j],s_negmodtile[j][i],s_negmodshouptile[j][i],mod); + } + + /* + if(i==0) + { + for(int j=0;j<8;j++) + { + printf("aa1 = %llu\n", (unsigned long long)Aa1[j]); + } + } + */ + + // ict Aa1 fix + + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + // modr Aa1 fix + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8]; + register uint64_t regs1[8]; + register uint64_t regs3[8]; + register uint64_t regs4[8]; + register uint64_t new_regs3[8]; + register uint64_t new_regs4[8]; + register uint64_t sum[8]; + register uint64_t prev_sum[8]; + register uint64_t tail[8]; + volatile __shared__ uint64_t carry_32[32]; + volatile __shared__ uint64_t sum_carry_32; + + volatile __shared__ uint64_t data_tail[256]; + + // register uint64_t fres[8]; + volatile __shared__ uint64_t carry; + // uint64_t carrySum=0; + volatile __shared__ uint64_t retail[2]; + + // register uint64_t regs_carry; + // new_regs3[0]=new_regs3[1]=0; + // new_regs4[0]=0; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + uint64_t temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); +// printf("carry32:%llu ",carry_32[i]); +// printf("after%2d: ",i); +#if 1 + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + + // mask = __activemask(); + // printf("线程 %2d: if条件满足,当前活跃掩码 = 0x%x\n", i, mask); + // printf("sum32:%llu ",sum_carry_32); + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whileAa1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + /* + for(int iter=0;iter<32;iter++) + { + int nextThread=i+1; + if(nextThread>1&&nextThread<=32) + { + tail[0]=tail[0]+carry_32[i]; + for(int j=0;j=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + //__syncwarp(); + } + } + */ + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + Aa1[0] = tail[0]; + Aa1[1] = tail[1]; + Aa1[2] = tail[2]; + Aa1[3] = tail[3]; + Aa1[4] = tail[4]; + Aa1[5] = tail[5]; + Aa1[6] = tail[6]; + Aa1[7] = tail[7]; + +#endif + + // ct Aa1 fix + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], twiddle[omegnIdx[0]], + twiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], twiddle[omegnIdx[1]], + twiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], twiddle[omegnIdx[2]], + twiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], twiddle[omegnIdx[3]], + twiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], twiddle[omegnIdx[4]], + twiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], twiddle[omegnIdx[5]], + twiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], twiddle[omegnIdx[6]], + twiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], twiddle[omegnIdx[7]], + twiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + // printf("regs[0]:%llu ",regs[0]); + } + + // GSbutter(Left,Right,mod,twiddle[4]); + // GSbutter(regs[0],regs[4],mod,twiddle[4]); + // GSbutter(regs[1],regs[5],mod,twiddle[5]); + // GSbutter(regs[2],regs[6],mod,twiddle[6]); + // GSbutter(regs[3],regs[7],mod,twiddle[7]); + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + // GSbutter(regs[0],regs[2],mod,twiddle[2]); + // GSbutter(regs[1],regs[3],mod,twiddle[3]); + // GSbutter(regs[4],regs[6],mod,twiddle[2]); + // GSbutter(regs[5],regs[7],mod,twiddle[3]); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + // GSbutter(regs[0],regs[1],mod,twiddle[1]); + // GSbutter(regs[2],regs[3],mod,twiddle[1]); + // GSbutter(regs[4],regs[5],mod,twiddle[1]); + // GSbutter(regs[6],regs[7],mod,twiddle[1]); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + Aa1[0] = regs[0]; + Aa1[1] = regs[1]; + Aa1[2] = regs[2]; + Aa1[3] = regs[3]; + Aa1[4] = regs[4]; + Aa1[5] = regs[5]; + Aa1[6] = regs[6]; + Aa1[7] = regs[7]; + + // ct(Aa1)*Rone fix + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(Aa1[j], CTinoutA[j], mod); + + // printf("M:%llu ",Mm1[j]); + } + // ict Mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * i + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * i + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * i + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * i + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * i + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * i + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * i + 7) % pairsInGroup); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + // GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + // GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + // GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + // GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + // GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + // GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + // GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + // GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0], right[0], ICTTwiddle[omegnIdx[0]], + ICTTwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], ICTTwiddle[omegnIdx[1]], + ICTTwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], ICTTwiddle[omegnIdx[2]], + ICTTwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], ICTTwiddle[omegnIdx[3]], + ICTTwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], ICTTwiddle[omegnIdx[4]], + ICTTwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], ICTTwiddle[omegnIdx[5]], + ICTTwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], ICTTwiddle[omegnIdx[6]], + ICTTwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], ICTTwiddle[omegnIdx[7]], + ICTTwiddle_shoup[omegnIdx[7]], mod); + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + // inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + + /* + //ict y fix + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + ct_butterfly(regs[0],regs[1],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[2],regs[3],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[4],regs[5],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + ct_butterfly(regs[6],regs[7],ICTTwiddle[1],ICTTwiddle_shoup[1],mod); + + ct_butterfly(regs[0],regs[2],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[1],regs[3],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + ct_butterfly(regs[4],regs[6],ICTTwiddle[2],ICTTwiddle_shoup[2],mod); + ct_butterfly(regs[5],regs[7],ICTTwiddle[3],ICTTwiddle_shoup[3],mod); + + ct_butterfly(regs[0],regs[4],ICTTwiddle[4],ICTTwiddle_shoup[4],mod); + ct_butterfly(regs[1],regs[5],ICTTwiddle[5],ICTTwiddle_shoup[5],mod); + ct_butterfly(regs[2],regs[6],ICTTwiddle[6],ICTTwiddle_shoup[6],mod); + ct_butterfly(regs[3],regs[7],ICTTwiddle[7],ICTTwiddle_shoup[7],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + } + + pairsInGroup=1<<(7); + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)%pairsInGroup+pairsInGroup; + omegnIdx[1]=pairsInGroup+((8*i+1)%pairsInGroup); + omegnIdx[2]=pairsInGroup+((8*i+2)%pairsInGroup); + omegnIdx[3]=pairsInGroup+((8*i+3)%pairsInGroup); + omegnIdx[4]=pairsInGroup+((8*i+4)%pairsInGroup); + omegnIdx[5]=pairsInGroup+((8*i+5)%pairsInGroup); + omegnIdx[6]=pairsInGroup+((8*i+6)%pairsInGroup); + omegnIdx[7]=pairsInGroup+((8*i+7)%pairsInGroup); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + //GSbutter(left[0],right[0],mod,twiddle[omegnIdx[0]]); + //GSbutter(left[1],right[1],mod,twiddle[omegnIdx[1]]); + //GSbutter(left[2],right[2],mod,twiddle[omegnIdx[2]]); + //GSbutter(left[3],right[3],mod,twiddle[omegnIdx[3]]); + //GSbutter(left[4],right[4],mod,twiddle[omegnIdx[4]]); + //GSbutter(left[5],right[5],mod,twiddle[omegnIdx[5]]); + //GSbutter(left[6],right[6],mod,twiddle[omegnIdx[6]]); + //GSbutter(left[7],right[7],mod,twiddle[omegnIdx[7]]); + ct_butterfly(left[0],right[0],ICTTwiddle[omegnIdx[0]],ICTTwiddle_shoup[omegnIdx[0]],mod); + ct_butterfly(left[1],right[1],ICTTwiddle[omegnIdx[1]],ICTTwiddle_shoup[omegnIdx[1]],mod); + ct_butterfly(left[2],right[2],ICTTwiddle[omegnIdx[2]],ICTTwiddle_shoup[omegnIdx[2]],mod); + ct_butterfly(left[3],right[3],ICTTwiddle[omegnIdx[3]],ICTTwiddle_shoup[omegnIdx[3]],mod); + ct_butterfly(left[4],right[4],ICTTwiddle[omegnIdx[4]],ICTTwiddle_shoup[omegnIdx[4]],mod); + ct_butterfly(left[5],right[5],ICTTwiddle[omegnIdx[5]],ICTTwiddle_shoup[omegnIdx[5]],mod); + ct_butterfly(left[6],right[6],ICTTwiddle[omegnIdx[6]],ICTTwiddle_shoup[omegnIdx[6]],mod); + ct_butterfly(left[7],right[7],ICTTwiddle[omegnIdx[7]],ICTTwiddle_shoup[omegnIdx[7]],mod); + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + for(int j=0;j<8;j++) + { + inoutA[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + //inout[8*i+j]=multiply_and_reduce_shoup_lazy(regs[j],inv,inv_shoup,mod); + } + */ + // modr Mm1 fix + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + if (i == 31) { + retail[0] = regs3[6] + regs4[7]; + // printf("retail[0]:%llu ",retail[0]); + retail[1] = regs3[7]; + // printf("retail[1]:%llu ",retail[1]); + } + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (i == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + // printf("sum7:%llu ",sum[7]); + carry = sum[7] >> u; + // printf("carry7:%llu ",carry); + } else { + for (int j = 0; j <= 7; j++) { + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + } + + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (i == 0) { + prev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } else { + prev_sum[0] = temp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + prev_sum[j] = sum[j - 1] >> u; + } + } + + tail[0] = (sum[0] & mod_mask) + prev_sum[0]; + tail[1] = (sum[1] & mod_mask) + prev_sum[1]; + tail[2] = (sum[2] & mod_mask) + prev_sum[2]; + tail[3] = (sum[3] & mod_mask) + prev_sum[3]; + tail[4] = (sum[4] & mod_mask) + prev_sum[4]; + tail[5] = (sum[5] & mod_mask) + prev_sum[5]; + tail[6] = (sum[6] & mod_mask) + prev_sum[6]; + tail[7] = (sum[7] & mod_mask) + prev_sum[7]; + // auto mask = __activemask(); + // printf("线程 %2d: a条件满足,当前活跃掩码 = 0x%x\n", i, mask); + + for (int j = 0; j <= 6; j++) { + if (tail[j] == 131072) { + tail[j] = 0; + tail[j + 1] += 1; + } + } + carry_32[i] = tail[7] >> u; + tail[7] = (tail[7] == 131072) ? 0 : (tail[7] & mod_mask); + + // carry_32[i]=tail[7]>>u; + // printf("before%2d: ",i); + __syncwarp(); + // printf("carry32:%llu ",carry_32[i]); + // printf("after%2d: ",i); + + // mask = __activemask(); + // printf("线程 %2d: 条件满足,当前活跃掩码 = 0x%x\n", i, mask); + //__syncwarp(__activemask()); + //__syncwarp(0x20); + // printf("carry32:%llu ",carry_32[i]); + + if (i == 0) { + sum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + sum_carry_32 += carry_32[j]; + } + } + + if (sum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + data_tail[8 * i + j] = tail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (data_tail[j] == 131072) { + data_tail[j + 1] += 1; + } + } + if (data_tail[255] == 131072) { + data_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + tail[j] = data_tail[8 * i + j]; + } + } + + /* + while(sum_carry_32!=0) + { + + tail[0]=(i==0)?tail[0]:(tail[0]+carry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(tail[j]==131072) + { + tail[j]=0; + tail[j+1]+=1; + } + } + carry_32[i]=tail[7]>>u; + tail[7]=(tail[7]==131072)?0:(tail[7]&mod_mask); + + if(i==0) + { + sum_carry_32=0; + for(int j=0;j<31;j++) + { + sum_carry_32+=carry_32[j]; + } + printf("sum32whilemm1!!!!!!!!:%llu ",sum_carry_32); + } + } + */ + + if (i == 0) { + uint64_t cay0 = (tail[0] + retail[0] + carry) >> u; + tail[0] = (tail[0] + retail[0] + carry) & mod_mask; + uint64_t cay1 = (tail[1] + retail[1] + cay0) >> u; + tail[1] = (tail[1] + retail[1] + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + + // mask = __activemask(); + // printf("线程 %2d: +2条件满足,当前活跃掩码 = 0x%x\n", i, mask); + } + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // nct x fix + + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // nct mm1 fix + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // nct y fix + + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + ((8 * i) / pair2); + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + ct_butterfly(left[0], right[0], NCTtwiddle[omegnIdx[0]], + NCTtwiddle_shoup[omegnIdx[0]], mod); + ct_butterfly(left[1], right[1], NCTtwiddle[omegnIdx[1]], + NCTtwiddle_shoup[omegnIdx[1]], mod); + ct_butterfly(left[2], right[2], NCTtwiddle[omegnIdx[2]], + NCTtwiddle_shoup[omegnIdx[2]], mod); + ct_butterfly(left[3], right[3], NCTtwiddle[omegnIdx[3]], + NCTtwiddle_shoup[omegnIdx[3]], mod); + ct_butterfly(left[4], right[4], NCTtwiddle[omegnIdx[4]], + NCTtwiddle_shoup[omegnIdx[4]], mod); + ct_butterfly(left[5], right[5], NCTtwiddle[omegnIdx[5]], + NCTtwiddle_shoup[omegnIdx[5]], mod); + ct_butterfly(left[6], right[6], NCTtwiddle[omegnIdx[6]], + NCTtwiddle_shoup[omegnIdx[6]], mod); + ct_butterfly(left[7], right[7], NCTtwiddle[omegnIdx[7]], + NCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[i + 32], NCTtwiddle_shoup[i + 32], + mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + (i * 2)], + NCTtwiddle_shoup[64 + (i * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + (i * 2) + 1], + NCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * i)], + NCTtwiddle_shoup[128 + (4 * i)], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * i) + 1], + NCTtwiddle_shoup[128 + (4 * i) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * i) + 2], + NCTtwiddle_shoup[128 + (4 * i) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * i) + 3], + NCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // nct(x)*rnct nct(y)*rnct nct(Mm1)*Nnct nct(Mm2)*Nnct g1+g2xandy + // nct x * nct y;nct(m)*nct(modn) + register uint64_t G1x[8]; + register uint64_t G2x[8]; + register uint64_t Gx[8]; + + for (int j = 0; j < 8; j++) { + G1x[j] = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + G2x[j] = multiply_and_reduce_shoup_lazy(Mm1[j], modn[8 * i + j], + modn_shoup[8 * i + j], mod); + Gx[j] = (G1x[j] + G2x[j]) % MOD; + } + + // inct gx fix + + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + // printf("inout[8*i]:%llu ",regs[0]); + // uint64_t nums=log2(n); + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * i)], + InvNCTtwiddle_shoup[128 + (4 * i)], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * i) + 1], + InvNCTtwiddle_shoup[128 + (4 * i) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * i) + 2], + InvNCTtwiddle_shoup[128 + (4 * i) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * i) + 3], + InvNCTtwiddle_shoup[128 + (4 * i) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + (i * 2)], + InvNCTtwiddle_shoup[64 + (i * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + (i * 2) + 1], + InvNCTtwiddle_shoup[64 + (i * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[i + 32], + InvNCTtwiddle_shoup[i + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> log_step) & 1; + // printf("C:%llu ",C); + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + } + + pairsInGroup = 1 << 7; + uint64_t pair2 = pairsInGroup << 1; + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + // printf("new_regs[0]:%llu ",new_regs[0]); + + C = (i >> 4) & 1; + // printf("C:%llu ",C); + + omegnIdx[0] = (8 * i) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + ((8 * i + 1) / pair2); + omegnIdx[2] = numOfGroups + ((8 * i + 2) / pair2); + omegnIdx[3] = numOfGroups + ((8 * i + 3) / pair2); + omegnIdx[4] = numOfGroups + ((8 * i + 4) / pair2); + omegnIdx[5] = numOfGroups + ((8 * i + 5) / pair2); + omegnIdx[6] = numOfGroups + ((8 * i + 6) / pair2); + omegnIdx[7] = numOfGroups + ((8 * i + 7) / pair2); + // printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0] = (1 - C) * (regs[0] - new_regs[0]) + new_regs[0]; + left[1] = (1 - C) * (regs[1] - new_regs[1]) + new_regs[1]; + left[2] = (1 - C) * (regs[2] - new_regs[2]) + new_regs[2]; + left[3] = (1 - C) * (regs[3] - new_regs[3]) + new_regs[3]; + left[4] = (1 - C) * (regs[4] - new_regs[4]) + new_regs[4]; + left[5] = (1 - C) * (regs[5] - new_regs[5]) + new_regs[5]; + left[6] = (1 - C) * (regs[6] - new_regs[6]) + new_regs[6]; + left[7] = (1 - C) * (regs[7] - new_regs[7]) + new_regs[7]; + + right[0] = C * (regs[0] - new_regs[0]) + new_regs[0]; + right[1] = C * (regs[1] - new_regs[1]) + new_regs[1]; + right[2] = C * (regs[2] - new_regs[2]) + new_regs[2]; + right[3] = C * (regs[3] - new_regs[3]) + new_regs[3]; + right[4] = C * (regs[4] - new_regs[4]) + new_regs[4]; + right[5] = C * (regs[5] - new_regs[5]) + new_regs[5]; + right[6] = C * (regs[6] - new_regs[6]) + new_regs[6]; + right[7] = C * (regs[7] - new_regs[7]) + new_regs[7]; + + gs_butterfly(left[0], right[0], InvNCTtwiddle[omegnIdx[0]], + InvNCTtwiddle_shoup[omegnIdx[0]], mod); + gs_butterfly(left[1], right[1], InvNCTtwiddle[omegnIdx[1]], + InvNCTtwiddle_shoup[omegnIdx[1]], mod); + gs_butterfly(left[2], right[2], InvNCTtwiddle[omegnIdx[2]], + InvNCTtwiddle_shoup[omegnIdx[2]], mod); + gs_butterfly(left[3], right[3], InvNCTtwiddle[omegnIdx[3]], + InvNCTtwiddle_shoup[omegnIdx[3]], mod); + gs_butterfly(left[4], right[4], InvNCTtwiddle[omegnIdx[4]], + InvNCTtwiddle_shoup[omegnIdx[4]], mod); + gs_butterfly(left[5], right[5], InvNCTtwiddle[omegnIdx[5]], + InvNCTtwiddle_shoup[omegnIdx[5]], mod); + gs_butterfly(left[6], right[6], InvNCTtwiddle[omegnIdx[6]], + InvNCTtwiddle_shoup[omegnIdx[6]], mod); + gs_butterfly(left[7], right[7], InvNCTtwiddle[omegnIdx[7]], + InvNCTtwiddle_shoup[omegnIdx[7]], mod); + + regs[0] = (1 - C) * (left[0] - right[0]) + right[0]; + regs[1] = (1 - C) * (left[1] - right[1]) + right[1]; + regs[2] = (1 - C) * (left[2] - right[2]) + right[2]; + regs[3] = (1 - C) * (left[3] - right[3]) + right[3]; + regs[4] = (1 - C) * (left[4] - right[4]) + right[4]; + regs[5] = (1 - C) * (left[5] - right[5]) + right[5]; + regs[6] = (1 - C) * (left[6] - right[6]) + right[6]; + regs[7] = (1 - C) * (left[7] - right[7]) + right[7]; + + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // inct y fix + /* + regs[0]=inoutA[8*i]; + regs[1]=inoutA[8*i+1]; + regs[2]=inoutA[8*i+2]; + regs[3]=inoutA[8*i+3]; + regs[4]=inoutA[8*i+4]; + regs[5]=inoutA[8*i+5]; + regs[6]=inoutA[8*i+6]; + regs[7]=inoutA[8*i+7]; + //printf("inout[8*i]:%llu ",regs[0]); + //uint64_t nums=log2(n); + gs_butterfly(regs[0],regs[1],InvNCTtwiddle[128+(4*i)],InvNCTtwiddle_shoup[128+(4*i)],mod); + gs_butterfly(regs[2],regs[3],InvNCTtwiddle[128+(4*i)+1],InvNCTtwiddle_shoup[128+(4*i)+1],mod); + gs_butterfly(regs[4],regs[5],InvNCTtwiddle[128+(4*i)+2],InvNCTtwiddle_shoup[128+(4*i)+2],mod); + gs_butterfly(regs[6],regs[7],InvNCTtwiddle[128+(4*i)+3],InvNCTtwiddle_shoup[128+(4*i)+3],mod); + + gs_butterfly(regs[0],regs[2],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[1],regs[3],InvNCTtwiddle[64+(i*2)],InvNCTtwiddle_shoup[64+(i*2)],mod); + gs_butterfly(regs[4],regs[6],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + gs_butterfly(regs[5],regs[7],InvNCTtwiddle[64+(i*2)+1],InvNCTtwiddle_shoup[64+(i*2)+1],mod); + + gs_butterfly(regs[0],regs[4],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[1],regs[5],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[2],regs[6],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + gs_butterfly(regs[3],regs[7],InvNCTtwiddle[i+32],InvNCTtwiddle_shoup[i+32],mod); + for(size_t log_m=4;log_m>=1;log_m--) + { + size_t log_step=4-log_m; + uint64_t pairsInGroup=1<<(log_step+3); + uint64_t numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<>log_step)&1; + //printf("C:%llu ",C); + uint64_t pair2=pairsInGroup<<1; + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + + + + + + } + + pairsInGroup=1<<7; + pair2=pairsInGroup<<1; + numOfGroups=128/pairsInGroup; + new_regs[0]=__shfl_xor_sync(0xffffffff,regs[0],1<<4); + new_regs[1]=__shfl_xor_sync(0xffffffff,regs[1],1<<4); + new_regs[2]=__shfl_xor_sync(0xffffffff,regs[2],1<<4); + new_regs[3]=__shfl_xor_sync(0xffffffff,regs[3],1<<4); + new_regs[4]=__shfl_xor_sync(0xffffffff,regs[4],1<<4); + new_regs[5]=__shfl_xor_sync(0xffffffff,regs[5],1<<4); + new_regs[6]=__shfl_xor_sync(0xffffffff,regs[6],1<<4); + new_regs[7]=__shfl_xor_sync(0xffffffff,regs[7],1<<4); + //printf("new_regs[0]:%llu ",new_regs[0]); + + C=(i>>4)&1; + //printf("C:%llu ",C); + + omegnIdx[0]=(8*i)/pair2+numOfGroups; + omegnIdx[1]=numOfGroups+((8*i+1)/pair2); + omegnIdx[2]=numOfGroups+((8*i+2)/pair2); + omegnIdx[3]=numOfGroups+((8*i+3)/pair2); + omegnIdx[4]=numOfGroups+((8*i+4)/pair2); + omegnIdx[5]=numOfGroups+((8*i+5)/pair2); + omegnIdx[6]=numOfGroups+((8*i+6)/pair2); + omegnIdx[7]=numOfGroups+((8*i+7)/pair2); + //printf("omegnIdx[0]:%llu ",omegnIdx[0]); + + left[0]=(1-C)*(regs[0]-new_regs[0])+new_regs[0]; + left[1]=(1-C)*(regs[1]-new_regs[1])+new_regs[1]; + left[2]=(1-C)*(regs[2]-new_regs[2])+new_regs[2]; + left[3]=(1-C)*(regs[3]-new_regs[3])+new_regs[3]; + left[4]=(1-C)*(regs[4]-new_regs[4])+new_regs[4]; + left[5]=(1-C)*(regs[5]-new_regs[5])+new_regs[5]; + left[6]=(1-C)*(regs[6]-new_regs[6])+new_regs[6]; + left[7]=(1-C)*(regs[7]-new_regs[7])+new_regs[7]; + + + + right[0]=C*(regs[0]-new_regs[0])+new_regs[0]; + right[1]=C*(regs[1]-new_regs[1])+new_regs[1]; + right[2]=C*(regs[2]-new_regs[2])+new_regs[2]; + right[3]=C*(regs[3]-new_regs[3])+new_regs[3]; + right[4]=C*(regs[4]-new_regs[4])+new_regs[4]; + right[5]=C*(regs[5]-new_regs[5])+new_regs[5]; + right[6]=C*(regs[6]-new_regs[6])+new_regs[6]; + right[7]=C*(regs[7]-new_regs[7])+new_regs[7]; + + + gs_butterfly(left[0],right[0],InvNCTtwiddle[omegnIdx[0]],InvNCTtwiddle_shoup[omegnIdx[0]],mod); + gs_butterfly(left[1],right[1],InvNCTtwiddle[omegnIdx[1]],InvNCTtwiddle_shoup[omegnIdx[1]],mod); + gs_butterfly(left[2],right[2],InvNCTtwiddle[omegnIdx[2]],InvNCTtwiddle_shoup[omegnIdx[2]],mod); + gs_butterfly(left[3],right[3],InvNCTtwiddle[omegnIdx[3]],InvNCTtwiddle_shoup[omegnIdx[3]],mod); + gs_butterfly(left[4],right[4],InvNCTtwiddle[omegnIdx[4]],InvNCTtwiddle_shoup[omegnIdx[4]],mod); + gs_butterfly(left[5],right[5],InvNCTtwiddle[omegnIdx[5]],InvNCTtwiddle_shoup[omegnIdx[5]],mod); + gs_butterfly(left[6],right[6],InvNCTtwiddle[omegnIdx[6]],InvNCTtwiddle_shoup[omegnIdx[6]],mod); + gs_butterfly(left[7],right[7],InvNCTtwiddle[omegnIdx[7]],InvNCTtwiddle_shoup[omegnIdx[7]],mod); + + + regs[0]=(1-C)*(left[0]-right[0])+right[0]; + regs[1]=(1-C)*(left[1]-right[1])+right[1]; + regs[2]=(1-C)*(left[2]-right[2])+right[2]; + regs[3]=(1-C)*(left[3]-right[3])+right[3]; + regs[4]=(1-C)*(left[4]-right[4])+right[4]; + regs[5]=(1-C)*(left[5]-right[5])+right[5]; + regs[6]=(1-C)*(left[6]-right[6])+right[6]; + regs[7]=(1-C)*(left[7]-right[7])+right[7]; + + inoutA[8*i]=multiply_and_reduce_shoup_lazy(regs[0],inv,inv_shoup,mod); + inoutA[8*i+1]=multiply_and_reduce_shoup_lazy(regs[1],inv,inv_shoup,mod); + inoutA[8*i+2]=multiply_and_reduce_shoup_lazy(regs[2],inv,inv_shoup,mod); + inoutA[8*i+3]=multiply_and_reduce_shoup_lazy(regs[3],inv,inv_shoup,mod); + inoutA[8*i+4]=multiply_and_reduce_shoup_lazy(regs[4],inv,inv_shoup,mod); + inoutA[8*i+5]=multiply_and_reduce_shoup_lazy(regs[5],inv,inv_shoup,mod); + inoutA[8*i+6]=multiply_and_reduce_shoup_lazy(regs[6],inv,inv_shoup,mod); + inoutA[8*i+7]=multiply_and_reduce_shoup_lazy(regs[7],inv,inv_shoup,mod); + */ + // modh Gx fix + // const uint64_t mod_mask=(1ULL< sample[8 * i]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[8 * i + 1]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[8 * i + 2]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[8 * i + 3]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[8 * i + 4]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[8 * i + 5]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[8 * i + 6]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[8 * i + 7]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + + sregs[0] = sregs0[0] >> u; + sregs[1] = sregs0[1] >> u; + sregs[2] = sregs0[2] >> u; + sregs[3] = sregs0[3] >> u; + sregs[4] = sregs0[4] >> u; + sregs[5] = sregs0[5] >> u; + sregs[6] = sregs0[6] >> u; + sregs[7] = sregs0[7] >> u; + + sregs1[0] = sregs0[0] & mod_mask; + sregs1[1] = sregs0[1] & mod_mask; + sregs1[2] = sregs0[2] & mod_mask; + sregs1[3] = sregs0[3] & mod_mask; + sregs1[4] = sregs0[4] & mod_mask; + sregs1[5] = sregs0[5] & mod_mask; + sregs1[6] = sregs0[6] & mod_mask; + sregs1[7] = sregs0[7] & mod_mask; + + sregs3[0] = sregs[0] >> u; + sregs3[1] = sregs[1] >> u; + sregs3[2] = sregs[2] >> u; + sregs3[3] = sregs[3] >> u; + sregs3[4] = sregs[4] >> u; + sregs3[5] = sregs[5] >> u; + sregs3[6] = sregs[6] >> u; + sregs3[7] = sregs[7] >> u; + + sregs4[0] = sregs[0] & mod_mask; + sregs4[1] = sregs[1] & mod_mask; + sregs4[2] = sregs[2] & mod_mask; + sregs4[3] = sregs[3] & mod_mask; + sregs4[4] = sregs[4] & mod_mask; + sregs4[5] = sregs[5] & mod_mask; + sregs4[6] = sregs[6] & mod_mask; + sregs4[7] = sregs[7] & mod_mask; + + if (i == 31) { + sretail[0] = sregs3[6] + sregs4[7]; + + sretail[1] = sregs3[7]; + } + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (i == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } else if (i == 31) { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry = ssum[7] >> u; + + } else { + for (int j = 0; j <= 7; j++) { + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + } + + uint64_t stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (i == 0) { + sprev_sum[0] = 0; // 线程0 显式置0,忽略 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } else { + sprev_sum[0] = stemp_shuffled; // 线程1-31 使用 shuffle 的结果 + for (int j = 1; j <= 7; j++) { + sprev_sum[j] = ssum[j - 1] >> u; + } + } + + stail[0] = (ssum[0] & mod_mask) + sprev_sum[0]; + stail[1] = (ssum[1] & mod_mask) + sprev_sum[1]; + stail[2] = (ssum[2] & mod_mask) + sprev_sum[2]; + stail[3] = (ssum[3] & mod_mask) + sprev_sum[3]; + stail[4] = (ssum[4] & mod_mask) + sprev_sum[4]; + stail[5] = (ssum[5] & mod_mask) + sprev_sum[5]; + stail[6] = (ssum[6] & mod_mask) + sprev_sum[6]; + stail[7] = (ssum[7] & mod_mask) + sprev_sum[7]; + + for (int j = 0; j <= 6; j++) { + if (stail[j] == 131072) { + stail[j] = 0; + stail[j + 1] += 1; + } + } + scarry_32[i] = stail[7] >> u; + stail[7] = (stail[7] == 131072) ? 0 : (stail[7] & mod_mask); + + __syncwarp(); + + if (i == 0) { + ssum_carry_32 = 0; + for (int j = 0; j < 31; j++) { + ssum_carry_32 += scarry_32[j]; + } + } + /* + while(ssum_carry_32!=0) + { + + stail[0]=(i==0)?stail[0]:(stail[0]+scarry_32[i-1]); + for(int j=0;j<=6;j++) + { + if(stail[j]==131072) + { + stail[j]=0; + stail[j+1]+=1; + } + } + scarry_32[i]=stail[7]>>u; + stail[7]=(stail[7]==131072)?0:(stail[7]&mod_mask); + + if(i==0) + { + ssum_carry_32=0; + for(int j=0;j<31;j++) + { + ssum_carry_32+=scarry_32[j]; + } + printf("sum32whilegx!!!!!!!!:%llu ",ssum_carry_32); + } + } + */ + if (ssum_carry_32 != 0) { + for (int j = 0; j < 8; j++) { + sdata_tail[8 * i + j] = stail[j]; + } + __syncwarp(); + if (i == 0) { + for (int j = 1; j <= 254; j++) { + if (sdata_tail[j] == 131072) { + sdata_tail[j + 1] += 1; + } + } + if (sdata_tail[255] == 131072) { + sdata_tail[255] = 0; + } + } + for (int j = 0; j < 8; j++) { + stail[j] = sdata_tail[8 * i + j]; + } + } + + if (i == 0) { + int64_t negcarry = 0; + + int64_t flag = stail[0] - sretail[0] - scarry; + + int64_t sum = (flag < 0) ? (flag + (1 << u)) : flag; + + negcarry = flag >> u; + + stail[0] = sum % (1 << u); + + flag = stail[1] - sretail[1] + negcarry; + sum = (flag < 0) ? (flag + (1 << u)) : flag; + negcarry = flag >> u; + + stail[1] = sum % (1 << u); + + stail[2] = (negcarry != 0) ? (stail[2] - 1) : stail[2]; + } + + uint64_t maxNum = (1 << u) - 1; + if (i == 0) { + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (i == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); + } +} + +// 【新增配置】每个 Block 允许的最大 Warp 数量 (m)。 +// 如果你每次调用是 128 线程 (m=4),可以把这里改为 +// 4,会极大节省编译期的共享内存分配! +#define MAX_WARPS 16 + +__global__ void XYfixWarpVector( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + // uint64_t warp_c, incoming_c; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + regs[0] = inoutA[8 * i]; + regs[1] = inoutA[8 * i + 1]; + regs[2] = inoutA[8 * i + 2]; + regs[3] = inoutA[8 * i + 3]; + regs[4] = inoutA[8 * i + 4]; + regs[5] = inoutA[8 * i + 5]; + regs[6] = inoutA[8 * i + 6]; + regs[7] = inoutA[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("gxmodh|%03d|%llu ", 8*lane_id+_j, (unsigned long long)stail[_j]); + printf("\n"); + __syncwarp(); + */ + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpROneVector( + uint64_t *inout, uint64_t *inoutAct, uint64_t *inoutAnct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + // uint64_t lc, warp_c, incoming_c; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + CTinoutA[0] = inoutAct[8 * lane_id + 0]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + /* + for(int _j = 0; _j < 8; _j++) + printf("m|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + /* + for(int _j = 0; _j < 8; _j++) + printf("mmodr|%03d|%llu ", 8*lane_id+_j, (unsigned long long)Mm1[_j]); + printf("\n"); + __syncwarp(); + */ + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + NCTinoutA[0] = inoutAnct[8 * lane_id + 0]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpSquareVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + + register uint64_t NCTinout[8]; + + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. ct(Aa1)*Rone fix vector ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) { + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + } + + // ================== 8. ict Mm1 fix ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + regs0[0] = Mm1[0] >> u; + regs0[1] = Mm1[1] >> u; + regs0[2] = Mm1[2] >> u; + regs0[3] = Mm1[3] >> u; + regs0[4] = Mm1[4] >> u; + regs0[5] = Mm1[5] >> u; + regs0[6] = Mm1[6] >> u; + regs0[7] = Mm1[7] >> u; + + regs1[0] = Mm1[0] & mod_mask; + regs1[1] = Mm1[1] & mod_mask; + regs1[2] = Mm1[2] & mod_mask; + regs1[3] = Mm1[3] & mod_mask; + regs1[4] = Mm1[4] & mod_mask; + regs1[5] = Mm1[5] & mod_mask; + regs1[6] = Mm1[6] & mod_mask; + regs1[7] = Mm1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +__global__ void XYfixWarpIRVector( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; + uint32_t lane_id = threadIdx.x & 0x1F; + // uint32_t warp_id = threadIdx.x >> 5; // 优化后已不再需要 shared + // memory,因此 warp_id 可以省略 + + register uint64_t CTinout[8]; + // register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + // register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + + register uint64_t C; + register uint32_t omegnIdx[8]; + + // 优化:提前声明并行进位需要的寄存器变量 + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + // ================== 1. CT inout fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone=ctinoutA ================== + + // ================== 3. x*negmodn fix vector ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) { + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + } + + // ================== 4. ict Aa1 fix vector ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup = 1 << (7); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = (8 * lane_id) % pairsInGroup + pairsInGroup; + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + + for (int j = 0; j < 8; j++) { + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + } + + // ================== 5. OPTIMIZED modr Aa1 fix vector ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + regs0[0] = Aa1[0] >> u; + regs0[1] = Aa1[1] >> u; + regs0[2] = Aa1[2] >> u; + regs0[3] = Aa1[3] >> u; + regs0[4] = Aa1[4] >> u; + regs0[5] = Aa1[5] >> u; + regs0[6] = Aa1[6] >> u; + regs0[7] = Aa1[7] >> u; + + regs1[0] = Aa1[0] & mod_mask; + regs1[1] = Aa1[1] & mod_mask; + regs1[2] = Aa1[2] & mod_mask; + regs1[3] = Aa1[3] & mod_mask; + regs1[4] = Aa1[4] & mod_mask; + regs1[5] = Aa1[5] & mod_mask; + regs1[6] = Aa1[6] & mod_mask; + regs1[7] = Aa1[7] & mod_mask; + + regs3[0] = regs0[0] >> u; + regs3[1] = regs0[1] >> u; + regs3[2] = regs0[2] >> u; + regs3[3] = regs0[3] >> u; + regs3[4] = regs0[4] >> u; + regs3[5] = regs0[5] >> u; + regs3[6] = regs0[6] >> u; + regs3[7] = regs0[7] >> u; + + regs4[0] = regs0[0] & mod_mask; + regs4[1] = regs0[1] & mod_mask; + regs4[2] = regs0[2] & mod_mask; + regs4[3] = regs0[3] & mod_mask; + regs4[4] = regs0[4] & mod_mask; + regs4[5] = regs0[5] & mod_mask; + regs4[6] = regs0[6] & mod_mask; + regs4[7] = regs0[7] & mod_mask; + + // 优化核心:使用 Shuffle 代替 Shared Memory 来传播 retail 溢出 + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = (regs1[1] + new_regs4[1]); + for (int j = 2; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + // 取代 carry[warp_id] + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j <= 7; j++) prev_sum[j] = sum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + } + + // ---------------------------------------------------- + // 第一阶段:寄存器内部本地进位 (消除 Shared Memory 循环) + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第一波从最高位(index 255)掉落的进位 + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + + /* + // 第二阶段:Warp 级并行前缀和 (Scan) 获取累计进位 + warp_c = lc; + #pragma unroll + for (int offset = 1; offset < 32; offset <<= 1) { + uint64_t t = __shfl_up_sync(0xffffffff, warp_c, offset); + if (lane_id >= offset) warp_c += t; + } + + // 传递累计进位给下一个线程 + incoming_c = __shfl_up_sync(0xffffffff, warp_c, 1); + if (lane_id == 0) incoming_c = 0; + + tail[0] += incoming_c; + + // 第三阶段:最终本地进位修正 + lc = 0; + #pragma unroll + for(int j=0; j<8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + + // 【修复】捕获第三阶段从最高位(index 255)再次掉落的进位 + uint64_t stage3_carry = __shfl_sync(0xffffffff, lc, 31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + // 第四阶段:Lane 0 单独处理边界回卷 + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + + tail[2] = (tail[2] + cay1) & mod_mask; + } + // ---------------------------------------------------- + Mm1[0] = tail[0]; + Mm1[1] = tail[1]; + Mm1[2] = tail[2]; + Mm1[3] = tail[3]; + Mm1[4] = tail[4]; + Mm1[5] = tail[5]; + Mm1[6] = tail[6]; + Mm1[7] = tail[7]; + + // ---------------------------------------------------- + + // ================== 6. ct Aa1 fix vector ================== + + // ================== 7. ct(Aa1)*Rone fix vector ================== + + // ================== 8. ict Mm1 fix ================== + + // ================== 9. OPTIMIZED modr Mm1 fix vector ================== + + // ================== 10. nct x fix vector ================== + regs[0] = inout[8 * i]; + regs[1] = inout[8 * i + 1]; + regs[2] = inout[8 * i + 2]; + regs[3] = inout[8 * i + 3]; + regs[4] = inout[8 * i + 4]; + regs[5] = inout[8 * i + 5]; + regs[6] = inout[8 * i + 6]; + regs[7] = inout[8 * i + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. nct mm1 fix vector ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + + ct_butterfly(regs[0], regs[4], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[(lane_id) + 32], + NCTtwiddle_shoup[(lane_id) + 32], mod); + + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + ((lane_id) * 2)], + NCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + ((lane_id) * 2) + 1], + NCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + (4 * (lane_id))], + NCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + (4 * (lane_id)) + 1], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + (4 * (lane_id)) + 2], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + (4 * (lane_id)) + 3], + NCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + Mm1[0] = regs[0]; + Mm1[1] = regs[1]; + Mm1[2] = regs[2]; + Mm1[3] = regs[3]; + Mm1[4] = regs[4]; + Mm1[5] = regs[5]; + Mm1[6] = regs[6]; + Mm1[7] = regs[7]; + + // ================== 12. nct y fix vector ================== + + // ================== 13. g1+g2xandy ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + // uint64_t tmp1 = + // multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j],NCTinout[j],mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. inct gx fix vector ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + (4 * (lane_id))], + InvNCTtwiddle_shoup[128 + (4 * (lane_id))], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + (4 * (lane_id)) + 1], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + (4 * (lane_id)) + 2], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + (4 * (lane_id)) + 3], + InvNCTtwiddle_shoup[128 + (4 * (lane_id)) + 3], mod); + + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + ((lane_id) * 2)], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2)], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + ((lane_id) * 2) + 1], + InvNCTtwiddle_shoup[64 + ((lane_id) * 2) + 1], mod); + + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[(lane_id) + 32], + InvNCTtwiddle_shoup[(lane_id) + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2 + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + + C = (lane_id >> 4) & 1; + + omegnIdx[0] = ((8 * lane_id) & 0xFF) / pair2_inv + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j]; + uint64_t r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + Gx[0] = multiply_and_reduce_shoup_lazy(regs[0], inv, inv_shoup, mod); + Gx[1] = multiply_and_reduce_shoup_lazy(regs[1], inv, inv_shoup, mod); + Gx[2] = multiply_and_reduce_shoup_lazy(regs[2], inv, inv_shoup, mod); + Gx[3] = multiply_and_reduce_shoup_lazy(regs[3], inv, inv_shoup, mod); + Gx[4] = multiply_and_reduce_shoup_lazy(regs[4], inv, inv_shoup, mod); + Gx[5] = multiply_and_reduce_shoup_lazy(regs[5], inv, inv_shoup, mod); + Gx[6] = multiply_and_reduce_shoup_lazy(regs[6], inv, inv_shoup, mod); + Gx[7] = multiply_and_reduce_shoup_lazy(regs[7], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp Shuffle + // 并行带符号修正) ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + // 最后写回外层大数组,使用全局索引 i + inout[8 * i] = static_cast(stail[0]); + inout[8 * i + 1] = static_cast(stail[1]); + inout[8 * i + 2] = static_cast(stail[2]); + inout[8 * i + 3] = static_cast(stail[3]); + inout[8 * i + 4] = static_cast(stail[4]); + inout[8 * i + 5] = static_cast(stail[5]); + inout[8 * i + 6] = static_cast(stail[6]); + inout[8 * i + 7] = static_cast(stail[7]); +} + +// start device function xy + +__device__ __forceinline__ void XYfixWarpVector_dev( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, uint64_t n, + uint64_t mod, const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. CT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinoutA[0] = regs[0]; + CTinoutA[1] = regs[1]; + CTinoutA[2] = regs[2]; + CTinoutA[3] = regs[3]; + CTinoutA[4] = regs[4]; + CTinoutA[5] = regs[5]; + CTinoutA[6] = regs[6]; + CTinoutA[7] = regs[7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. NCT inoutA ================== + // ★ inoutA[8*i+j] → inoutA[8*lane_id+j] + regs[0] = inoutA[8 * lane_id]; + regs[1] = inoutA[8 * lane_id + 1]; + regs[2] = inoutA[8 * lane_id + 2]; + regs[3] = inoutA[8 * lane_id + 3]; + regs[4] = inoutA[8 * lane_id + 4]; + regs[5] = inoutA[8 * lane_id + 5]; + regs[6] = inoutA[8 * lane_id + 6]; + regs[7] = inoutA[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinoutA[0] = regs[0]; + NCTinoutA[1] = regs[1]; + NCTinoutA[2] = regs[2]; + NCTinoutA[3] = regs[3]; + NCTinoutA[4] = regs[4]; + NCTinoutA[5] = regs[5]; + NCTinoutA[6] = regs[6]; + NCTinoutA[7] = regs[7]; + + // ================== 13. Gx = NCT(x)*NCT(y) + NCT(m)*N ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device function xy + +// start device xr1 2banben +__device__ __forceinline__ void XYfixWarpROneVector_dev( + uint64_t *inout, + uint64_t *inoutAct, // 预计算好的 CT(R1),所有 warp 共享,不偏移 + uint64_t *inoutAnct, // 预计算好的 NCT(R1),所有 warp 共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t CTinoutA[8]; + register uint64_t NCTinout[8]; + register uint64_t NCTinoutA[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. ctRone = + // CTinoutA(预计算,直接读,不偏移)================== + CTinoutA[0] = inoutAct[8 * lane_id]; + CTinoutA[1] = inoutAct[8 * lane_id + 1]; + CTinoutA[2] = inoutAct[8 * lane_id + 2]; + CTinoutA[3] = inoutAct[8 * lane_id + 3]; + CTinoutA[4] = inoutAct[8 * lane_id + 4]; + CTinoutA[5] = inoutAct[8 * lane_id + 5]; + CTinoutA[6] = inoutAct[8 * lane_id + 6]; + CTinoutA[7] = inoutAct[8 * lane_id + 7]; + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinoutA ================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinoutA[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. nctRone = + // NCTinoutA(预计算,直接读,不偏移)================== + NCTinoutA[0] = inoutAnct[8 * lane_id]; + NCTinoutA[1] = inoutAnct[8 * lane_id + 1]; + NCTinoutA[2] = inoutAnct[8 * lane_id + 2]; + NCTinoutA[3] = inoutAnct[8 * lane_id + 3]; + NCTinoutA[4] = inoutAnct[8 * lane_id + 4]; + NCTinoutA[5] = inoutAnct[8 * lane_id + 5]; + NCTinoutA[6] = inoutAnct[8 * lane_id + 6]; + NCTinoutA[7] = inoutAnct[8 * lane_id + 7]; + + // ================== 13. Gx ================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinoutA[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xr1 2baneben + +// start device xx +__device__ __forceinline__ void XYfixWarpSquareVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 为空(平方不需要独立 + // CTinoutA)================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // ================== 6. CT Aa1 ================== + regs[0] = tail[0]; + regs[1] = tail[1]; + regs[2] = tail[2]; + regs[3] = tail[3]; + regs[4] = tail[4]; + regs[5] = tail[5]; + regs[6] = tail[6]; + regs[7] = tail[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + + // ================== 7. CT(Aa1)*CTinout(平方:y=x,直接用 + // CTinout)================== + register uint64_t Mm1[8]; + for (int j = 0; j < 8; j++) + Mm1[j] = + multiply_and_reduce_shoup_lazy_dotproduct(regs[j], CTinout[j], mod); + + // ================== 8. ICT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Mm1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 9. modr Mm1 ================== + for (int j = 0; j < 8; j++) regs0[j] = Mm1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Mm1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + accumulated_wrap = 0; + + active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 为空(平方不需要独立 + // NCTinoutA)================== + + // ================== 13. Gx = NCT(x)*NCT(x) + + // NCT(m)*N(平方:y=x)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp1 = multiply_and_reduce_shoup_lazy_dotproduct(NCTinout[j], + NCTinout[j], mod); + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (tmp1 + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. modh Gx ================== + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device xx + +// start device x1 +__device__ __forceinline__ void XYfixWarpIRVector_dev( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n, uint64_t mod, + const __restrict__ uint64_t *twiddle, + const __restrict__ uint64_t *twiddle_shoup, + const __restrict__ uint64_t *ICTTwiddle, + const __restrict__ uint64_t *ICTTwiddle_shoup, + const __restrict__ uint64_t *NCTtwiddle, + const __restrict__ uint64_t *NCTtwiddle_shoup, + const __restrict__ uint64_t *InvNCTtwiddle, + const __restrict__ uint64_t *InvNCTtwiddle_shoup, + const __restrict__ uint64_t inv, const __restrict__ uint64_t inv_shoup, + const __restrict__ uint64_t *sample) { + // ★ 唯一修改:去掉全局索引 i,改用 lane_id + uint32_t lane_id = threadIdx.x & 0x1F; + + register uint64_t CTinout[8]; + register uint64_t NCTinout[8]; + register uint64_t regs[8]; + register uint64_t new_regs[8]; + register uint64_t C; + register uint32_t omegnIdx[8]; + + uint64_t retail_0, retail_1, carry_wrap, temp_shuffled; + uint64_t lc; + int64_t sretail_0, sretail_1, scarry_wrap, stemp_shuffled; + int64_t slc; + // int64_t swarp_c, sincoming_c; + + // ================== 1. CT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, twiddle[omegnIdx[j]], twiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + gs_butterfly(regs[0], regs[4], twiddle[4], twiddle_shoup[4], mod); + gs_butterfly(regs[1], regs[5], twiddle[5], twiddle_shoup[5], mod); + gs_butterfly(regs[2], regs[6], twiddle[6], twiddle_shoup[6], mod); + gs_butterfly(regs[3], regs[7], twiddle[7], twiddle_shoup[7], mod); + gs_butterfly(regs[0], regs[2], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[1], regs[3], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[4], regs[6], twiddle[2], twiddle_shoup[2], mod); + gs_butterfly(regs[5], regs[7], twiddle[3], twiddle_shoup[3], mod); + gs_butterfly(regs[0], regs[1], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[2], regs[3], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[4], regs[5], twiddle[1], twiddle_shoup[1], mod); + gs_butterfly(regs[6], regs[7], twiddle[1], twiddle_shoup[1], mod); + CTinout[0] = regs[0]; + CTinout[1] = regs[1]; + CTinout[2] = regs[2]; + CTinout[3] = regs[3]; + CTinout[4] = regs[4]; + CTinout[5] = regs[5]; + CTinout[6] = regs[6]; + CTinout[7] = regs[7]; + + // ================== 2. Section 2 + // 为空(IR:y=1,无需单独CT(y))================== + + // ================== 3. x*negmodn ================== + register uint64_t Aa1[8]; + register uint64_t Mm1[8]; + + for (int j = 0; j < 8; j++) + Aa1[j] = + multiply_and_reduce_shoup_lazy(CTinout[j], negmodn[8 * lane_id + j], + negmodn_shoup[8 * lane_id + j], mod); + + // ================== 4. ICT Aa1 ================== + regs[0] = Aa1[0]; + regs[1] = Aa1[1]; + regs[2] = Aa1[2]; + regs[3] = Aa1[3]; + regs[4] = Aa1[4]; + regs[5] = Aa1[5]; + regs[6] = Aa1[6]; + regs[7] = Aa1[7]; + ct_butterfly(regs[0], regs[1], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[2], regs[3], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[4], regs[5], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[6], regs[7], ICTTwiddle[1], ICTTwiddle_shoup[1], mod); + ct_butterfly(regs[0], regs[2], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[1], regs[3], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[4], regs[6], ICTTwiddle[2], ICTTwiddle_shoup[2], mod); + ct_butterfly(regs[5], regs[7], ICTTwiddle[3], ICTTwiddle_shoup[3], mod); + ct_butterfly(regs[0], regs[4], ICTTwiddle[4], ICTTwiddle_shoup[4], mod); + ct_butterfly(regs[1], regs[5], ICTTwiddle[5], ICTTwiddle_shoup[5], mod); + ct_butterfly(regs[2], regs[6], ICTTwiddle[6], ICTTwiddle_shoup[6], mod); + ct_butterfly(regs[3], regs[7], ICTTwiddle[7], ICTTwiddle_shoup[7], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup = 1 << 7; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = pairsInGroup + ((8 * lane_id) % pairsInGroup); + omegnIdx[1] = pairsInGroup + ((8 * lane_id + 1) % pairsInGroup); + omegnIdx[2] = pairsInGroup + ((8 * lane_id + 2) % pairsInGroup); + omegnIdx[3] = pairsInGroup + ((8 * lane_id + 3) % pairsInGroup); + omegnIdx[4] = pairsInGroup + ((8 * lane_id + 4) % pairsInGroup); + omegnIdx[5] = pairsInGroup + ((8 * lane_id + 5) % pairsInGroup); + omegnIdx[6] = pairsInGroup + ((8 * lane_id + 6) % pairsInGroup); + omegnIdx[7] = pairsInGroup + ((8 * lane_id + 7) % pairsInGroup); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, ICTTwiddle[omegnIdx[j]], ICTTwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Aa1[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 5. modr Aa1 ================== + int u = 17; + const uint64_t mod_mask = (1ULL << u) - 1; + register uint64_t regs0[8], regs1[8], regs3[8], regs4[8]; + register uint64_t new_regs3[8], new_regs4[8]; + register uint64_t sum[8], prev_sum[8], tail[8]; + + for (int j = 0; j < 8; j++) regs0[j] = Aa1[j] >> u; + for (int j = 0; j < 8; j++) regs1[j] = Aa1[j] & mod_mask; + for (int j = 0; j < 8; j++) regs3[j] = regs0[j] >> u; + for (int j = 0; j < 8; j++) regs4[j] = regs0[j] & mod_mask; + + retail_0 = __shfl_sync(0xffffffff, regs3[6] + regs4[7], 31); + retail_1 = __shfl_sync(0xffffffff, regs3[7], 31); + new_regs3[0] = __shfl_up_sync(0xffffffff, regs3[6], 1); + new_regs3[1] = __shfl_up_sync(0xffffffff, regs3[7], 1); + new_regs3[2] = regs3[0]; + new_regs3[3] = regs3[1]; + new_regs3[4] = regs3[2]; + new_regs3[5] = regs3[3]; + new_regs3[6] = regs3[4]; + new_regs3[7] = regs3[5]; + new_regs4[0] = __shfl_up_sync(0xffffffff, regs4[7], 1); + new_regs4[1] = regs4[0]; + new_regs4[2] = regs4[1]; + new_regs4[3] = regs4[2]; + new_regs4[4] = regs4[3]; + new_regs4[5] = regs4[4]; + new_regs4[6] = regs4[5]; + new_regs4[7] = regs4[6]; + + if (lane_id == 0) { + sum[0] = regs1[0]; + sum[1] = regs1[1] + new_regs4[1]; + for (int j = 2; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } else { + for (int j = 0; j < 8; j++) sum[j] = new_regs3[j] + regs1[j] + new_regs4[j]; + } + + carry_wrap = __shfl_sync(0xffffffff, sum[7] >> u, 31); + temp_shuffled = __shfl_up_sync(0xffffffff, sum[7] >> u, 1); + if (lane_id == 0) { + prev_sum[0] = 0; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } else { + prev_sum[0] = temp_shuffled; + for (int j = 1; j < 8; j++) prev_sum[j] = sum[j - 1] >> u; + } +#pragma unroll + for (int j = 0; j < 8; j++) tail[j] = (sum[j] & mod_mask) + prev_sum[j]; + + lc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += lc; + lc = tail[j] >> u; + tail[j] &= mod_mask; + } + uint64_t stage1_carry = __shfl_sync(0xffffffff, lc, 31); + /* + warp_c=lc; + #pragma unroll + for(int offset=1;offset<32;offset<<=1){ + uint64_t t=__shfl_up_sync(0xffffffff,warp_c,offset); + if(lane_id>=(uint32_t)offset) warp_c+=t; + } + incoming_c=__shfl_up_sync(0xffffffff,warp_c,1); + if(lane_id==0) incoming_c=0; + tail[0]+=incoming_c; + + lc=0; + #pragma unroll + for(int j=0;j<8;j++){ tail[j]+=lc; lc=tail[j]>>u; tail[j]&=mod_mask; } + uint64_t stage3_carry=__shfl_sync(0xffffffff,lc,31); + */ + uint64_t accumulated_wrap = 0; + + uint32_t active = __ballot_sync(0xffffffff, lc != 0); + while (active != 0) { + // 收集 lane 31 溢出(用于模运算回卷) + if (lane_id == 31) accumulated_wrap += lc; + + // 只把进位传给紧邻的下一个 lane + uint64_t incoming = __shfl_up_sync(0xffffffff, lc, 1); + if (lane_id == 0) incoming = 0; // 回卷单独处理 + + lc = 0; + if (incoming != 0) { +#pragma unroll + for (int j = 0; j < 8; j++) { + tail[j] += incoming; + incoming = tail[j] >> u; + tail[j] &= mod_mask; + } + lc = incoming; + } + active = __ballot_sync(0xffffffff, lc != 0); + } + + // 取 lane 31 累积的回卷进位 + uint64_t final_wrap = __shfl_sync(0xffffffff, accumulated_wrap, 31); + + if (lane_id == 0) { + // uint64_t total_wrap0 = retail_0 + carry_wrap + stage1_carry + + // stage3_carry; + uint64_t total_wrap0 = retail_0 + carry_wrap + final_wrap; + uint64_t cay0 = (tail[0] + total_wrap0) >> u; + tail[0] = (tail[0] + total_wrap0) & mod_mask; + uint64_t cay1 = (tail[1] + retail_1 + cay0) >> u; + tail[1] = (tail[1] + retail_1 + cay0) & mod_mask; + tail[2] = (tail[2] + cay1) & mod_mask; + } + + // modr 结果存入 Mm1(Section 6/7/8/9 为空,IR直接用 Aa1 的 modr 结果) + for (int j = 0; j < 8; j++) Mm1[j] = tail[j]; + + // ================== 10. NCT inout ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + regs[0] = inout[8 * lane_id]; + regs[1] = inout[8 * lane_id + 1]; + regs[2] = inout[8 * lane_id + 2]; + regs[3] = inout[8 * lane_id + 3]; + regs[4] = inout[8 * lane_id + 4]; + regs[5] = inout[8 * lane_id + 5]; + regs[6] = inout[8 * lane_id + 6]; + regs[7] = inout[8 * lane_id + 7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + NCTinout[0] = regs[0]; + NCTinout[1] = regs[1]; + NCTinout[2] = regs[2]; + NCTinout[3] = regs[3]; + NCTinout[4] = regs[4]; + NCTinout[5] = regs[5]; + NCTinout[6] = regs[6]; + NCTinout[7] = regs[7]; + + // ================== 11. NCT Mm1 ================== + regs[0] = Mm1[0]; + regs[1] = Mm1[1]; + regs[2] = Mm1[2]; + regs[3] = Mm1[3]; + regs[4] = Mm1[4]; + regs[5] = Mm1[5]; + regs[6] = Mm1[6]; + regs[7] = Mm1[7]; + for (size_t log_m = 0; log_m < 5; log_m++) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = numOfGroups + (((8 * lane_id) & 0xFF) / pair2); + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + ct_butterfly(l, r, NCTtwiddle[omegnIdx[j]], NCTtwiddle_shoup[omegnIdx[j]], + mod); + regs[j] = C ? r : l; + } + } + ct_butterfly(regs[0], regs[4], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[1], regs[5], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[2], regs[6], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[3], regs[7], NCTtwiddle[lane_id + 32], + NCTtwiddle_shoup[lane_id + 32], mod); + ct_butterfly(regs[0], regs[2], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[1], regs[3], NCTtwiddle[64 + lane_id * 2], + NCTtwiddle_shoup[64 + lane_id * 2], mod); + ct_butterfly(regs[4], regs[6], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[5], regs[7], NCTtwiddle[64 + lane_id * 2 + 1], + NCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + ct_butterfly(regs[0], regs[1], NCTtwiddle[128 + 4 * lane_id], + NCTtwiddle_shoup[128 + 4 * lane_id], mod); + ct_butterfly(regs[2], regs[3], NCTtwiddle[128 + 4 * lane_id + 1], + NCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + ct_butterfly(regs[4], regs[5], NCTtwiddle[128 + 4 * lane_id + 2], + NCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + ct_butterfly(regs[6], regs[7], NCTtwiddle[128 + 4 * lane_id + 3], + NCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + for (int j = 0; j < 8; j++) Mm1[j] = regs[j]; + + // ================== 12. Section 12 + // 为空(IR:y=1,NCT(1)直接就是1)================== + + // ================== 13. Gx = NCT(x)*1 + + // NCT(m)*N(IR:NCTinoutA=1)================== + register uint64_t Gx[8]; + for (int j = 0; j < 8; j++) { + uint64_t tmp2 = multiply_and_reduce_shoup_lazy( + Mm1[j], modn[8 * lane_id + j], modn_shoup[8 * lane_id + j], mod); + Gx[j] = (NCTinout[j] + tmp2) % MOD; + } + + // ================== 14. INCT Gx ================== + regs[0] = Gx[0]; + regs[1] = Gx[1]; + regs[2] = Gx[2]; + regs[3] = Gx[3]; + regs[4] = Gx[4]; + regs[5] = Gx[5]; + regs[6] = Gx[6]; + regs[7] = Gx[7]; + gs_butterfly(regs[0], regs[1], InvNCTtwiddle[128 + 4 * lane_id], + InvNCTtwiddle_shoup[128 + 4 * lane_id], mod); + gs_butterfly(regs[2], regs[3], InvNCTtwiddle[128 + 4 * lane_id + 1], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 1], mod); + gs_butterfly(regs[4], regs[5], InvNCTtwiddle[128 + 4 * lane_id + 2], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 2], mod); + gs_butterfly(regs[6], regs[7], InvNCTtwiddle[128 + 4 * lane_id + 3], + InvNCTtwiddle_shoup[128 + 4 * lane_id + 3], mod); + gs_butterfly(regs[0], regs[2], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[1], regs[3], InvNCTtwiddle[64 + lane_id * 2], + InvNCTtwiddle_shoup[64 + lane_id * 2], mod); + gs_butterfly(regs[4], regs[6], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[5], regs[7], InvNCTtwiddle[64 + lane_id * 2 + 1], + InvNCTtwiddle_shoup[64 + lane_id * 2 + 1], mod); + gs_butterfly(regs[0], regs[4], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[1], regs[5], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[2], regs[6], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + gs_butterfly(regs[3], regs[7], InvNCTtwiddle[lane_id + 32], + InvNCTtwiddle_shoup[lane_id + 32], mod); + for (size_t log_m = 4; log_m >= 1; log_m--) { + size_t log_step = 4 - log_m; + uint64_t pairsInGroup = 1 << (log_step + 3); + uint64_t numOfGroups = 128 / pairsInGroup; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << log_step); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << log_step); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << log_step); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << log_step); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << log_step); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << log_step); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << log_step); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << log_step); + C = (lane_id >> log_step) & 1; + uint64_t pair2 = pairsInGroup << 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2) + numOfGroups; + omegnIdx[1] = numOfGroups + (((8 * lane_id + 1) & 0xFF) / pair2); + omegnIdx[2] = numOfGroups + (((8 * lane_id + 2) & 0xFF) / pair2); + omegnIdx[3] = numOfGroups + (((8 * lane_id + 3) & 0xFF) / pair2); + omegnIdx[4] = numOfGroups + (((8 * lane_id + 4) & 0xFF) / pair2); + omegnIdx[5] = numOfGroups + (((8 * lane_id + 5) & 0xFF) / pair2); + omegnIdx[6] = numOfGroups + (((8 * lane_id + 6) & 0xFF) / pair2); + omegnIdx[7] = numOfGroups + (((8 * lane_id + 7) & 0xFF) / pair2); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + { + uint64_t pairsInGroup_inv = 1 << 7; + uint64_t pair2_inv = pairsInGroup_inv << 1; + uint64_t numOfGroups_inv = 128 / pairsInGroup_inv; + new_regs[0] = __shfl_xor_sync(0xffffffff, regs[0], 1 << 4); + new_regs[1] = __shfl_xor_sync(0xffffffff, regs[1], 1 << 4); + new_regs[2] = __shfl_xor_sync(0xffffffff, regs[2], 1 << 4); + new_regs[3] = __shfl_xor_sync(0xffffffff, regs[3], 1 << 4); + new_regs[4] = __shfl_xor_sync(0xffffffff, regs[4], 1 << 4); + new_regs[5] = __shfl_xor_sync(0xffffffff, regs[5], 1 << 4); + new_regs[6] = __shfl_xor_sync(0xffffffff, regs[6], 1 << 4); + new_regs[7] = __shfl_xor_sync(0xffffffff, regs[7], 1 << 4); + C = (lane_id >> 4) & 1; + omegnIdx[0] = (((8 * lane_id) & 0xFF) / pair2_inv) + numOfGroups_inv; + omegnIdx[1] = numOfGroups_inv + (((8 * lane_id + 1) & 0xFF) / pair2_inv); + omegnIdx[2] = numOfGroups_inv + (((8 * lane_id + 2) & 0xFF) / pair2_inv); + omegnIdx[3] = numOfGroups_inv + (((8 * lane_id + 3) & 0xFF) / pair2_inv); + omegnIdx[4] = numOfGroups_inv + (((8 * lane_id + 4) & 0xFF) / pair2_inv); + omegnIdx[5] = numOfGroups_inv + (((8 * lane_id + 5) & 0xFF) / pair2_inv); + omegnIdx[6] = numOfGroups_inv + (((8 * lane_id + 6) & 0xFF) / pair2_inv); + omegnIdx[7] = numOfGroups_inv + (((8 * lane_id + 7) & 0xFF) / pair2_inv); + for (int j = 0; j < 8; j++) { + uint64_t l = C ? new_regs[j] : regs[j], r = C ? regs[j] : new_regs[j]; + gs_butterfly(l, r, InvNCTtwiddle[omegnIdx[j]], + InvNCTtwiddle_shoup[omegnIdx[j]], mod); + regs[j] = C ? r : l; + } + } + for (int j = 0; j < 8; j++) + Gx[j] = multiply_and_reduce_shoup_lazy(regs[j], inv, inv_shoup, mod); + + // ================== 15. OPTIMIZED modh Gx fix vector (Warp 涟漪进位修正) + // ================== + + register int64_t sregs0[8], sregs[8], sregs1[8], sregs3[8], sregs4[8]; + register int64_t snew_regs3[8], snew_regs4[8]; + register int64_t ssum[8], sprev_sum[8], stail[8]; + register uint64_t DivMod[8]; + + // 1. 获取带符号的初始值 + sregs0[0] = (Gx[0] > sample[8 * lane_id]) + ? (static_cast(Gx[0]) - static_cast(MOD)) + : static_cast(Gx[0]); + sregs0[1] = (Gx[1] > sample[(8 * lane_id + 1) & 0xFF]) + ? (static_cast(Gx[1]) - static_cast(MOD)) + : static_cast(Gx[1]); + sregs0[2] = (Gx[2] > sample[(8 * lane_id + 2) & 0xFF]) + ? (static_cast(Gx[2]) - static_cast(MOD)) + : static_cast(Gx[2]); + sregs0[3] = (Gx[3] > sample[(8 * lane_id + 3) & 0xFF]) + ? (static_cast(Gx[3]) - static_cast(MOD)) + : static_cast(Gx[3]); + sregs0[4] = (Gx[4] > sample[(8 * lane_id + 4) & 0xFF]) + ? (static_cast(Gx[4]) - static_cast(MOD)) + : static_cast(Gx[4]); + sregs0[5] = (Gx[5] > sample[(8 * lane_id + 5) & 0xFF]) + ? (static_cast(Gx[5]) - static_cast(MOD)) + : static_cast(Gx[5]); + sregs0[6] = (Gx[6] > sample[(8 * lane_id + 6) & 0xFF]) + ? (static_cast(Gx[6]) - static_cast(MOD)) + : static_cast(Gx[6]); + sregs0[7] = (Gx[7] > sample[(8 * lane_id + 7) & 0xFF]) + ? (static_cast(Gx[7]) - static_cast(MOD)) + : static_cast(Gx[7]); + +// 利用算术右移自动保留符号位 +#pragma unroll + for (int j = 0; j < 8; j++) { + sregs[j] = sregs0[j] >> u; + sregs1[j] = sregs0[j] & mod_mask; + sregs3[j] = sregs[j] >> u; + sregs4[j] = sregs[j] & mod_mask; + } + + sretail_0 = __shfl_sync(0xffffffff, sregs3[6] + sregs4[7], 31); + sretail_1 = __shfl_sync(0xffffffff, sregs3[7], 31); + + snew_regs3[0] = __shfl_up_sync(0xffffffff, sregs3[6], 1); + snew_regs3[1] = __shfl_up_sync(0xffffffff, sregs3[7], 1); + snew_regs3[2] = sregs3[0]; + snew_regs3[3] = sregs3[1]; + snew_regs3[4] = sregs3[2]; + snew_regs3[5] = sregs3[3]; + snew_regs3[6] = sregs3[4]; + snew_regs3[7] = sregs3[5]; + + snew_regs4[0] = __shfl_up_sync(0xffffffff, sregs4[7], 1); + snew_regs4[1] = sregs4[0]; + snew_regs4[2] = sregs4[1]; + snew_regs4[3] = sregs4[2]; + snew_regs4[4] = sregs4[3]; + snew_regs4[5] = sregs4[4]; + snew_regs4[6] = sregs4[5]; + snew_regs4[7] = sregs4[6]; + + if (lane_id == 0) { + ssum[0] = sregs1[0]; + ssum[1] = (sregs1[1] + snew_regs4[1]); + for (int j = 2; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } else { + for (int j = 0; j <= 7; j++) + ssum[j] = snew_regs3[j] + sregs1[j] + snew_regs4[j]; + } + + scarry_wrap = __shfl_sync(0xffffffff, ssum[7] >> u, 31); + stemp_shuffled = __shfl_up_sync(0xffffffff, ssum[7] >> u, 1); + + if (lane_id == 0) { + sprev_sum[0] = 0; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } else { + sprev_sum[0] = stemp_shuffled; + for (int j = 1; j <= 7; j++) sprev_sum[j] = ssum[j - 1] >> u; + } + +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] = (ssum[j] & mod_mask) + sprev_sum[j]; + } + + // ========================================================================= + // ★ 核心修复:Stage 1 & 2 使用 Ballot While 循环涟漪进位,替代错误的前缀和 + // ========================================================================= + slc = 0; +#pragma unroll + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + + int64_t accumulated_lane31_carry = + 0; // 专门捕捉所有从最高位(lane 31)溢出的最终进位/借位 + active = __ballot_sync(0xffffffff, slc != 0); + + // 只要有任何一个线程还有进位/借位没消化,整个 Warp 就继续同步传递 + while (active != 0) { + if (lane_id == 31) { + accumulated_lane31_carry += slc; + } + + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; // 最低位暂时不回卷 + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; // 算术右移完美处理负数借位 + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // 获取从最高位最终溢出/借位的所有总量 + int64_t final_carry = __shfl_sync(0xffffffff, accumulated_lane31_carry, 31); + + // ========================================================================= + // ★ 核心修复:Stage 3 (Mod H 负向回卷) + // ========================================================================= + slc = 0; + if (lane_id == 0) { + // 对于模 h (2^4352 + 1),最高位的进位溢出权重大于 0,等价于从最低位减去。 + // 总溢出 = sretail_0 + scarry_wrap + final_carry + int64_t total_wrap_0 = sretail_0 + scarry_wrap + final_carry; + + int64_t flag = stail[0] - total_wrap_0; + int64_t negcarry = flag >> u; + stail[0] = flag & mod_mask; + + flag = stail[1] - sretail_1 + negcarry; + negcarry = flag >> u; + stail[1] = flag & mod_mask; + + // 让借位顺着 lane_id 0 往上走,走到没有为止 + for (int j = 2; j < 8; j++) { + flag = stail[j] + negcarry; + negcarry = flag >> u; + stail[j] = flag & mod_mask; + } + + slc = negcarry; // 如果 lane 0 扣完 8 个元素后还欠借位,扔给后面的线程 + } + + // 再次开启涟漪循环,彻底消化回卷带来的多米诺借位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + // ========================================================================= + // Stage 4: 计算最终的 w = (h - g) / 2 + // ========================================================================= + uint64_t maxNum = (1ULL << u) - 1; + if (lane_id == 0) { + // 这一步等价于取补码 + 2 + stail[0] = maxNum + 2 - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } else { + stail[0] = maxNum - stail[0]; + stail[1] = maxNum - stail[1]; + stail[2] = maxNum - stail[2]; + stail[3] = maxNum - stail[3]; + stail[4] = maxNum - stail[4]; + stail[5] = maxNum - stail[5]; + stail[6] = maxNum - stail[6]; + stail[7] = maxNum - stail[7]; + } + + // ★ 修复漏洞:maxNum+2-stail[0] 如果 stail[0] 极小,会导致超过 17 + // 位限制产生进位 + slc = 0; + if (lane_id == 0) { + for (int j = 0; j < 8; j++) { + stail[j] += slc; + slc = stail[j] >> u; + stail[j] &= mod_mask; + } + } + // 传播这最后一丝极小的正向进位 + active = __ballot_sync(0xffffffff, slc != 0); + while (active != 0) { + int64_t incoming_c = __shfl_up_sync(0xffffffff, slc, 1); + if (lane_id == 0) incoming_c = 0; + slc = 0; + if (incoming_c != 0) { + for (int j = 0; j < 8; j++) { + stail[j] += incoming_c; + incoming_c = stail[j] >> u; + stail[j] &= mod_mask; + } + slc = incoming_c; + } + active = __ballot_sync(0xffffffff, slc != 0); + } + + uint64_t temp_shuffled1 = __shfl_down_sync(0xffffffff, stail[0], 1); + if (lane_id == 31) { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = 0; + } else { + DivMod[0] = (stail[1] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[1] = (stail[2] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[2] = (stail[3] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[3] = (stail[4] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[4] = (stail[5] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[5] = (stail[6] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[6] = (stail[7] % 2 == 1) ? (1 << (u - 1)) : 0; + DivMod[7] = (temp_shuffled1 % 2 == 1) ? (1 << (u - 1)) : 0; + } + + stail[0] = (stail[0] / 2) + DivMod[0]; + stail[1] = (stail[1] / 2) + DivMod[1]; + stail[2] = (stail[2] / 2) + DivMod[2]; + stail[3] = (stail[3] / 2) + DivMod[3]; + stail[4] = (stail[4] / 2) + DivMod[4]; + stail[5] = (stail[5] / 2) + DivMod[5]; + stail[6] = (stail[6] / 2) + DivMod[6]; + stail[7] = (stail[7] / 2) + DivMod[7]; + + // ================== 16. 写回 ================== + // ★ inout[8*i+j] → inout[8*lane_id+j] + for (int j = 0; j < 8; j++) + inout[8 * lane_id + j] = static_cast(stail[j]); +} + +// end device x1 +//------------------------test------------------------------------- +// test device xy +__global__ void XYfixWarpVector_dev_wrapper( + uint64_t *inout, uint64_t *inoutA, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + // 每个 warp 在全局中的编号 + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 把指针偏移到本 warp 专属的 256 元素段 + uint64_t *my_inout = inout + global_warp_id * 256; + uint64_t *my_inoutA = inoutA + global_warp_id * 256; + + XYfixWarpVector_dev(my_inout, my_inoutA, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); +} + +// test device xr1 +// ================================================================ +// 包装核函数:调用设备函数版本 +// 启动参数与原核函数完全相同:<<>> +// ================================================================ +__global__ void XYfixWarpROneVector_dev_wrapper( + uint64_t *inout, + uint64_t *inoutAct, // CT(y),256元素,所有warp共享,不偏移 + uint64_t *inoutAnct, // NCT(y),256元素,所有warp共享,不偏移 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // 只有 inout 按 warp 偏移 + uint64_t *my_inout = inout + global_warp_id * 256; + // inoutAct / inoutAnct 不偏移,所有 warp 共享同一份 256 元素 + uint64_t *my_inoutAct = inoutAct; + uint64_t *my_inoutAnct = inoutAnct; + + XYfixWarpROneVector_dev(my_inout, my_inoutAct, my_inoutAnct, negmodn, + negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpSquareVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpSquareVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test xx device end + +// test x1 start +// ================================================================ +// 包装核函数(验证用):启动参数 <<>> 与原核函数一致 +// ================================================================ +__global__ void XYfixWarpIRVector_dev_wrapper( + uint64_t *inout, uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warps_per_block = blockDim.x >> 5; + const int warp_in_block = threadIdx.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + + // inout 按 warp 偏移,其余参数全部共享 + uint64_t *my_inout = inout + global_warp_id * 256; + + XYfixWarpIRVector_dev( + my_inout, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); +} + +// test x1 end + +// ================================================================ +// FMLE 核函数(Algorithm 5: FFT-Based McLaughlin's Exponentiation) +// +// 共享内存布局(每 warp 512 个 uint64_t = 4096 bytes): +// ws[0..255] : y_buf(标准形式) +// ws[256..511] : t_buf(标准形式) +// +// 预计算参数(主机端准备好后传入,大小均为 256,所有 warp 共享): +// d_r0 : r₀ = r mod n +// d_r1_ct : CT(r₁),其中 r₁ = r² mod n +// d_r1_nct : NCT(r₁) +// +// 指数表示:exp_bits[i] = e[i](0 或 1),i=0 为 LSB,i=tau-1 为 MSB +// +// +// ================================================================ +__global__ void FMLEKernel( + const uint64_t *__restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t *__restrict__ exp_bits, // [tau] 指数位,e[0]=LSB + int tau, // 指数位数 + uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t *__restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; // [256] y + uint64_t *const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移(与之前验证一致) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, // ← 全局共享,不偏移 + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // exp_bits[i] 对 warp 内所有 lane 相同 → 无 warp 分歧 + // ================================================================ + for (int i = 0; i < tau; i++) { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + // 使用 XYfixWarpVector_dev:inout=t_buf(读写),inoutA=y_buf(只读) + // t_buf 和 y_buf 是不同缓冲区,无别名问题 + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, // t = t × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + // 使用 XYfixWarpSquareVector_dev:原地平方 + XYfixWarpSquareVector_dev(y_buf, // y = y × y + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, + mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 使用 XYfixWarpIRVector_dev:去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev(t_buf, // t = t × 1(去 Montgomery) + negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, + twiddle, twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, InvNCTtwiddle, + InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: return t → 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void FMLEKernel_Debug( + const uint64_t *__restrict__ A, const uint64_t *__restrict__ exp_bits, + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample, + // ---- 新增调试参数 ---- + uint64_t *debug_out, // [6 * 256] 存储6个中间快照 + int debug_warp_id) // 要调试的 warp 编号 +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 是否是需要调试的 warp + const bool is_debug = (global_warp_id == debug_warp_id); + +// 快照宏:把 buf[0..255] 写入 debug_out 的第 slot 槽 +#define SNAPSHOT(slot, buf) \ + if (is_debug) { \ + for (int _j = 0; _j < 8; _j++) \ + debug_out[(slot) * 256 + 8 * lane_id + _j] = (buf)[8 * lane_id + _j]; \ + __syncwarp(); \ + } + +// ==================== Step 6: y = x ==================== +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(0, y_buf) // slot 0: y after Step6 + +// ==================== Step 7: t = r0 ==================== +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + SNAPSHOT(1, t_buf) // slot 1: t after Step7 + + // ==================== Step 8: y = FMLM(y, r1) ==================== + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(2, y_buf) // slot 2: y after Step8 + + // ==================== Loop: Steps 10 & 11 ==================== + for (int i = 0; i < tau; i++) { + // Step 10: t = FMLM(t, y) + if (exp_bits[i] == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // 只记录前4轮的 Step10 结果(slot 3 固定记录最后一次 t) + } + SNAPSHOT(3, t_buf) + // Step 11: y = FMLM(y, y) + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + + if (i < 4) { + SNAPSHOT(4, y_buf) // slot 4: y after Step11(第 i 轮) + } + } + + // ==================== Step 13: t = FMLM(t, 1) ==================== + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + + SNAPSHOT(5, t_buf) // slot 5: t after Step13(最终结果) + +#undef SNAPSHOT + +// 写回最终输出 +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup); + +__global__ void Testnctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup) { + __shared__ uint64_t buffer[256]; + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i; + + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = 1; numOfGroups <= n / 2; + numOfGroups = numOfGroups << 1) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + if (numOfGroups == 1) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // CTbutter(samples[0],samples[1],mod,omgn); + ct_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + + if (numOfGroups == (n / 2)) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + + __syncthreads(); // 线程同步 + } + } + } +} + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup); + +__global__ void Testinctsample(uint64_t *inout, uint64_t n, uint64_t mod, + uint64_t *twiddle, uint64_t *twiddle_shoup, + uint64_t inv, uint64_t inv_shoup) { + __shared__ uint64_t buffer[256]; + // uint64_t inv=ksm(n,mod-2,mod); + for (uint64_t i = blockIdx.x * blockDim.x + threadIdx.x; i < n / 2; + i += blockDim.x * gridDim.x) { + uint64_t tid = i % (n / 2); + uint64_t pairsInGroup; + uint64_t k, j, glbIdx, glbIdxadd; // k = psi_step + uint64_t samples[2]; + uint64_t omgn, omgn_shoup; + uint64_t twiddleId; + + for (uint64_t numOfGroups = n / 2; numOfGroups >= 1; + numOfGroups = numOfGroups = numOfGroups / 2) { + pairsInGroup = n / numOfGroups / 2; + + k = tid / pairsInGroup; // numOfGroups是这一轮NTT变换中的分组数 + j = tid % pairsInGroup; + // pairInGroup 是每个分组中多少对; k=psi_step? k是组号;j是对号 + glbIdx = 2 * k * pairsInGroup + j; + glbIdxadd = glbIdx + pairsInGroup; // 计算全局内存的索引 + + // twiddleId=pairsInGroup+j; + twiddleId = k + numOfGroups; + omgn = twiddle[twiddleId]; + omgn_shoup = twiddle_shoup[twiddleId]; + + // printf("num:%llu ",numOfGroups); + // printf("j:%llu ",j); + // printf("tid:%llu ",tid); + + // printf("twiddleid:%llu ",twiddleId); + // printf("pair:%llu ",pairsInGroup); + // printf("glbIdx:%llu ",glbIdx); + // printf("omgn:%llu ",omgn); + // printf("glbIdxadd:%llu ",glbIdxadd); + // uuint64_t psi_shoup = twiddles_shoup[numOfGroups + k + n * mod_idx]; + + // printf("sample[glb]:%llu ",samples[0]); + // printf("sample[glbIdxadd]:%llu ",samples[1]); + if (numOfGroups == (n / 2)) { + samples[0] = inout[glbIdx]; + samples[1] = inout[glbIdxadd]; + + } else { + samples[0] = buffer[glbIdx]; + samples[1] = buffer[glbIdxadd]; + } + + // GSbutter(samples[0],samples[1],mod,omgn); + gs_butterfly(samples[0], samples[1], omgn, omgn_shoup, mod); + // printf("Asample[glb]:%llu ",samples[0]); + // printf("Asample[glbIdxadd]:%llu ",samples[1]); + + if (numOfGroups == 1) { + inout[glbIdx] = samples[0]; + + inout[glbIdxadd] = samples[1]; + inout[glbIdx] = + multiply_and_reduce_shoup_lazy(inout[glbIdx], inv, inv_shoup, mod); + inout[glbIdxadd] = multiply_and_reduce_shoup_lazy(inout[glbIdxadd], inv, + inv_shoup, mod); + // inout[glbIdx]=ksc(inout[glbIdx],inv,mod); + // inout[glbIdxadd]=ksc(inout[glbIdxadd],inv,mod); + + } else { + buffer[glbIdx] = samples[0]; + + buffer[glbIdxadd] = samples[1]; + __syncthreads(); // 线程同步 + // printf("inout[glb]:%llu ",inout[glbIdx]); + // printf("inout[glbIdxadd]:%llu ",inout[glbIdxadd]); + } + } + } +} + +struct CommonDeviceArrays { + uint64_t *d_negmodn, *d_negmodn_shoup; + uint64_t *d_modn, *d_modn_shoup; + uint64_t *d_twiddle, *d_twiddle_shoup; + uint64_t *d_ICTTwiddle, *d_ICTTwiddle_shoup; + uint64_t *d_NCTtwiddle, *d_NCTtwiddle_shoup; + uint64_t *d_InvNCTtwiddle, *d_InvNCTtwiddle_shoup; + uint64_t *d_sample; + uint64_t n; + uint64_t mod; + uint64_t inv; + uint64_t inv_shoup; +}; + +void multiStream(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + CommonDeviceArrays &common) { + const int num_elements_per_stream = 256; + size_t chunk_bytes = num_elements_per_stream * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 分配设备端的完整 arrX 和 arrY + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t streams[128]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 分发任务:H2D拷贝 -> 内核执行 -> D2H拷贝 + for (int i = 0; i < n_streams; i++) { + int offset = i * num_elements_per_stream; + + // 异步拷贝数据到 GPU + // cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + // cudaMemcpyHostToDevice, streams[i]); cudaMemcpyAsync(&d_arrY[offset], + // &arrY[offset], chunk_bytes, cudaMemcpyHostToDevice, streams[i]); + cudaMemcpy(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice); + // 启动内核 + XYfixWarp<<<1, 32>>>(&d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, + common.d_modn_shoup, common.n, common.mod, + common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, + common.d_NCTtwiddle, common.d_NCTtwiddle_shoup, + common.d_InvNCTtwiddle, common.d_InvNCTtwiddle_shoup, + common.inv, common.inv_shoup, common.d_sample); + + // 异步将结果拷回 Host + // cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + // cudaMemcpyDeviceToHost, streams[i]); cudaMemcpyAsync(&arrY[offset], + // &d_arrY[offset], chunk_bytes, cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpy(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + cudaMemcpy(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost); + } + + // 同步所有流并清理 + cudaDeviceSynchronize(); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVectorTime(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // ========================================== + // 引入 CUDA Event 进行高精度异步计时 + // ========================================== + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // 记录 H2D 拷贝开始 + cudaEventRecord(start_h2d, 0); // 0 表示默认的空流 (串行控制点) + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + + // 记录 H2D 结束 + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 记录 Kernel 执行开始 + cudaEventRecord(start_ker, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + + // 记录 Kernel 执行结束 + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 记录 D2H 拷贝开始 + cudaEventRecord(start_d2h, 0); + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + // 记录 D2H 拷贝结束 + cudaEventRecord(end_d2h, 0); + // ========================================== + + // 阻塞 CPU,等待 GPU 把所有流的任务以及 Event 打点全部做完 + cudaDeviceSynchronize(); + + // ========================================== + // 提取并计算 Event 记录的物理时间 + // ========================================== + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [深度性能剖析 (基于 CUDA Event)] ===" << std::endl; + // cudaEventElapsedTime 默认返回毫秒(ms),我们乘 1000 转为微秒(us) + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理资源 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void multiStreamVector(uint64_t *arrX, uint64_t *arrY, size_t n_streams, + size_t n_blocks, CommonDeviceArrays &common) { + const int LIMBS = 256; + + // 每个流现在要处理 n_blocks * 256 个元素 + size_t chunk_elements = n_blocks * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 动态创建 n_streams 个 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 异步拷贝 H2D + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + + // 启动内核:现在使用的是 <<>> + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + + // 异步拷贝 D2H + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + + cudaDeviceSynchronize(); + + // 清理资源 + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +void generate_random_array(uint64_t arr[128]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 128; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +void generate_random_array256(uint64_t arr[256]) { + // 初始化随机数种子 + std::srand(static_cast(std::time(nullptr))); + + for (int i = 0; i < 256; ++i) { + // 生成0-32767之间的随机数,然后扩展到0-131071 + int r1 = std::rand() & 0x7FFF; // 0-32767 + int r2 = std::rand() & 0x7FFF; // 0-32767 + int r3 = std::rand() & 0x7FFF; // 0-32767 + int r4 = std::rand() & 0x7FFF; // 0-32767 + + // 组合成一个0-131071之间的数 + int combined = ((r1 + r2 + r3 + r4) * 131071) / (4 * 32767); + + // 确保不超出范围 + if (combined > 131071) combined = 131071; + arr[i] = static_cast(combined); + } +} + +// test multiwarp end + +void multiStream_m_warps_Time(uint64_t *arrX, uint64_t *arrY, int n_streams, + int poly_per_stream, int m, + CommonDeviceArrays &common) { + const int LIMBS = 256; + + // ========================================== + // 核心维度计算 (升级点) + // ========================================== + // 1. 数据量计算:每个流处理的多项式个数 * 256 + size_t chunk_elements = poly_per_stream * LIMBS; + size_t chunk_bytes = chunk_elements * sizeof(uint64_t); + size_t total_bytes = n_streams * chunk_bytes; + + // 2. 线程与网格维度计算: + // 因为 1 个 Block 能吃掉 m 个多项式,所以需要的 Block 数量要除以 m (向上取整) + int grid_size = (poly_per_stream + m - 1) / m; + int threads_per_block = m * 32; + + std::cout << "\n[GPU 调度信息] 流数量: " << n_streams + << ", 每流多项式数: " << poly_per_stream << "\n[Kernel 参数] <<< " + << grid_size << " Blocks, " << threads_per_block + << " Threads >>> per stream\n"; + + // ========================================== + // 分配设备端的大显存池 + // ========================================== + uint64_t *d_arrX, *d_arrY; + cudaMalloc(&d_arrX, total_bytes); + cudaMalloc(&d_arrY, total_bytes); + + // 创建 CUDA 流 + cudaStream_t *streams = new cudaStream_t[n_streams]; + for (int i = 0; i < n_streams; i++) { + cudaStreamCreate(&streams[i]); + } + + // 引入 CUDA Event 进行高精度异步计时 + cudaEvent_t start_h2d, end_h2d; + cudaEvent_t start_ker, end_ker; + cudaEvent_t start_d2h, end_d2h; + + cudaEventCreate(&start_h2d); + cudaEventCreate(&end_h2d); + cudaEventCreate(&start_ker); + cudaEventCreate(&end_ker); + cudaEventCreate(&start_d2h); + cudaEventCreate(&end_d2h); + + // ------------------------------------------ + // 1. 记录 H2D 拷贝 + cudaEventRecord(start_h2d, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&d_arrX[offset], &arrX[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + cudaMemcpyAsync(&d_arrY[offset], &arrY[offset], chunk_bytes, + cudaMemcpyHostToDevice, streams[i]); + } + cudaEventRecord(end_h2d, 0); + + // ------------------------------------------ + // 2. 记录 Kernel 执行 + cudaEventRecord(start_ker, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + + // 启动内核,精准传入重新计算好的 grid_size 和 threads_per_block + XYfixWarpVector<<>>( + &d_arrX[offset], &d_arrY[offset], common.d_negmodn, + common.d_negmodn_shoup, common.d_modn, common.d_modn_shoup, common.n, + common.mod, common.d_twiddle, common.d_twiddle_shoup, + common.d_ICTTwiddle, common.d_ICTTwiddle_shoup, common.d_NCTtwiddle, + common.d_NCTtwiddle_shoup, common.d_InvNCTtwiddle, + common.d_InvNCTtwiddle_shoup, common.inv, common.inv_shoup, + common.d_sample); + } + cudaEventRecord(end_ker, 0); + + // ------------------------------------------ + // 3. 记录 D2H 拷贝 + cudaEventRecord(start_d2h, 0); + for (int i = 0; i < n_streams; i++) { + int offset = i * chunk_elements; + cudaMemcpyAsync(&arrX[offset], &d_arrX[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + cudaMemcpyAsync(&arrY[offset], &d_arrY[offset], chunk_bytes, + cudaMemcpyDeviceToHost, streams[i]); + } + cudaEventRecord(end_d2h, 0); + + // ========================================== + // 阻塞 CPU,等待所有异步操作落锤 + // ========================================== + cudaDeviceSynchronize(); + + float ms_h2d = 0, ms_ker = 0, ms_d2h = 0; + cudaEventElapsedTime(&ms_h2d, start_h2d, end_h2d); + cudaEventElapsedTime(&ms_ker, start_ker, end_ker); + cudaEventElapsedTime(&ms_d2h, start_d2h, end_d2h); + + std::cout << "\n=== [m-Warp 多流并发性能剖析] ===" << std::endl; + std::cout << "1. Host -> Device 拷贝耗时 : " << ms_h2d * 1000.0f << " us" + << std::endl; + std::cout << "2. Kernel(s) 并发执行耗时 : " << ms_ker * 1000.0f << " us" + << std::endl; + std::cout << "3. Device -> Host 拷贝耗时 : " << ms_d2h * 1000.0f << " us" + << std::endl; + std::cout << "=========================================\n" << std::endl; + + // 清理大扫除 + cudaEventDestroy(start_h2d); + cudaEventDestroy(end_h2d); + cudaEventDestroy(start_ker); + cudaEventDestroy(end_ker); + cudaEventDestroy(start_d2h); + cudaEventDestroy(end_d2h); + + for (int i = 0; i < n_streams; i++) { + cudaStreamDestroy(streams[i]); + } + delete[] streams; + + // 只释放函数内部申请的局部缓存 + cudaFree(d_arrX); + cudaFree(d_arrY); +} + +#define NLIMBS 256 +#define LIMB_BITS 17 +#define MASK_17 ((1ULL << 17) - 1) + +#define M_CASES 512 // 总测试用例数 +#define N_STREAMS 4 // 并发流数 +#define N_BLOCKS 8 // 每次内核启动的 block 数 +#define WARPS_PER_BLOCK 8 // 每 block 的 warp 数(16×32=512 线程 ≤ 1024) + +// 每次内核调用处理的用例数 +#define CASES_PER_LAUNCH (N_BLOCKS * WARPS_PER_BLOCK) + +// 每流负责的用例数 +#define CASES_PER_STREAM ((M_CASES + N_STREAMS - 1) / N_STREAMS) + +// 每流需要的内核调用批次数 +#define BATCHES_PER_STREAM \ + ((CASES_PER_STREAM + CASES_PER_LAUNCH - 1) / CASES_PER_LAUNCH) + +// --------------------------------------------------------------------------- +// CUDA 错误检查 +// --------------------------------------------------------------------------- +#define CUDA_CHECK(call) \ + do { \ + cudaError_t _e = (call); \ + if (_e != cudaSuccess) { \ + fprintf(stderr, "[CUDA Error] %s:%d %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(_e)); \ + exit(EXIT_FAILURE); \ + } \ + } while (0) + +// ============================================================================= +// 内核调用宏(带流,grid = n_blocks, block = warps_per_block * 32) +// ============================================================================= +#define CALL_XY_S(_nb, _nt, _io, _ioA, _st) \ + XYfixWarpVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), (_ioA), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, \ + 256, MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_R1_S(_nb, _nt, _io, _st) \ + XYfixWarpROneVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_r1_ct, d_r1_nct, d_negModn, d_con_NegModn_shoup, d_Modn, \ + d_con_Modn_shoup, 256, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample) + +#define CALL_SQ_S(_nb, _nt, _io, _st) \ + XYfixWarpSquareVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +#define CALL_IR_S(_nb, _nt, _io, _st) \ + XYfixWarpIRVector<<<(_nb), (_nt), 0, (_st)>>>( \ + (_io), d_negModn, d_con_NegModn_shoup, d_Modn, d_con_Modn_shoup, 256, \ + MOD, d_con_twiddle, d_con_twiddle_shoup, d_con_ICTTwiddle, \ + d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup, \ + d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, inv_shoup, d_sample) + +// ============================================================================= +// StreamResource — 每个流的独立资源 +// ============================================================================= +struct StreamResource { + cudaStream_t stream; + uint64_t *d_x; // 设备端底数 + uint64_t *d_result; // 设备端结果 + uint64_t *d_y; // 工作寄存器 y(整流全部用例) + uint64_t *d_t; // 工作寄存器 t + uint64_t *h_x_pinned; // 锁页主机底数缓冲 + uint64_t *h_result_pinned; // 锁页主机结果缓冲 + int cases_in_stream; +}; + +// ============================================================================= +// FMLE_stream_batch() +// 在指定流上,用 n_blocks 个 block(每 block WARPS_PER_BLOCK 个 warp) +// 同时处理 batch_size = n_blocks × WARPS_PER_BLOCK 个用例。 +// +// 内核启动维度: +// grid = (n_blocks, 1, 1) +// block = (WARPS_PER_BLOCK * 32, 1, 1) +// +// 内核内部用全局 warp ID 索引数据: +// global_warp_id = blockIdx.x * WARPS_PER_BLOCK + threadIdx.x / 32 +// data_offset = global_warp_id * 256 +// ============================================================================= +static void FMLE_stream_batch( + uint64_t *d_result_batch, uint64_t *d_x_batch, uint64_t *d_y_buf, + uint64_t *d_t_buf, const uint64_t *e_limbs, int tau_limbs, + int n_blocks, // 本次实际使用的 block 数 + int warps_per_block, cudaStream_t stream, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, uint64_t *d_con_NegModn_shoup, + uint64_t *d_Modn, uint64_t *d_con_Modn_shoup, uint64_t MOD, + uint64_t *d_con_twiddle, uint64_t *d_con_twiddle_shoup, + uint64_t *d_con_ICTTwiddle, uint64_t *d_con_ICTTwiddle_shoup, + uint64_t *d_con_twiddle_NCT, uint64_t *d_con_twiddle_NCT_shoup, + uint64_t *d_con_InvTwiddle, uint64_t *d_con_InvTwiddle_shoup, uint64_t inv, + uint64_t inv_shoup, uint64_t *d_sample) { + // grid = n_blocks 个 block,每 block = warps_per_block * 32 线程 + const int threads_per_block = warps_per_block * 32; // ≤ 1024 + const int batch_size = n_blocks * warps_per_block; + const size_t sz = (size_t)batch_size * NLIMBS * sizeof(uint64_t); + + // 步骤 6: y ← x + CUDA_CHECK(cudaMemcpyAsync(d_y_buf, d_x_batch, sz, cudaMemcpyDeviceToDevice, + stream)); + + // 步骤 7: t ← r0(广播到 batch_size 个槽位) + for (int p = 0; p < batch_size; p++) { + CUDA_CHECK(cudaMemcpyAsync(d_t_buf + (size_t)p * NLIMBS, d_r0, + NLIMBS * sizeof(uint64_t), + cudaMemcpyDeviceToDevice, stream)); + } + + // 步骤 8: y ← FMLM(y, r1) + CALL_R1_S(n_blocks, threads_per_block, d_y_buf, stream); + + // 步骤 9-12: 主循环 + const int tau_bits = tau_limbs * LIMB_BITS; + for (int i = 0; i < tau_bits; i++) { + const int bit = (int)((e_limbs[i / LIMB_BITS] >> (i % LIMB_BITS)) & 1ULL); + + // 步骤 10 + if (bit == 1) + CALL_XY_S(n_blocks, threads_per_block, d_t_buf, d_y_buf, stream); + + // 步骤 11 + CALL_SQ_S(n_blocks, threads_per_block, d_y_buf, stream); + } + + // 步骤 13: t ← FMLM(t, 1) + CALL_IR_S(n_blocks, threads_per_block, d_t_buf, stream); + + // 步骤 14: 写出结果 + CUDA_CHECK(cudaMemcpyAsync(d_result_batch, d_t_buf, sz, + cudaMemcpyDeviceToDevice, stream)); +} + +// ============================================================================= +// writeLimbsLine — 写 limb 数组到文件(单行) +// ============================================================================= +static void writeLimbsLine(FILE *fp, const uint64_t *a, int len) { + for (int i = 0; i < len; i++) { + fprintf(fp, "%llu", (unsigned long long)a[i]); + if (i < len - 1) fprintf(fp, ","); + } + fprintf(fp, "\n"); +} + +/* +__global__ void FMLMKernel_e( + const uint64_t * __restrict__ A, // [m*256] 输入 x(标准形式) + const uint64_t * __restrict__ exp_bits, // [m*tau] +指数位,每对一段,e[i*tau+0]=LSB int tau, // +指数位数 uint64_t *output, // [m*256] 输出 t = x^e mod n + int m, // 大数个数 + // ---- 预计算常量(大小 256,所有 warp 共享,不按 warp 偏移)---- + const uint64_t * __restrict__ d_r0, // r₀ = r mod n + uint64_t *d_r1_ct, // CT(r₁),r₁ = r² mod n + uint64_t *d_r1_nct, // NCT(r₁) + // ---- 模乘运算参数 ---- + uint64_t *negmodn, uint64_t *negmodn_shoup, + uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) +{ + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + +warp_in_block; const uint32_t lane_id = threadIdx.x & 0x1F; + + if(global_warp_id >= m) return; + + // ---- 共享内存分区:每 warp 独占 512 个 uint64_t ---- + extern __shared__ uint64_t smem[]; + uint64_t * const ws = smem + warp_in_block * 512; + uint64_t * const y_buf = ws; // [256] y + uint64_t * const t_buf = ws + 256; // [256] t + + const int off = global_warp_id * 256; + + // ================================================================ + // Step 6: y ← x + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + y_buf[8*lane_id + j] = A[off + 8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 7: t ← r₀ + // d_r0 是全局共享常量(256个元素),所有 warp 读同一份,不偏移 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + t_buf[8*lane_id + j] = d_r0[8*lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // 使用 XYfixWarpROneVector_dev + // d_r1_ct / d_r1_nct 是全局共享,不按 warp 偏移 + // ================================================================ + XYfixWarpROneVector_dev( + y_buf, + d_r1_ct, d_r1_nct, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // 每个 warp 读取自己对应的指数段:exp_bits + global_warp_id * tau + // ================================================================ + const uint64_t *my_exp = exp_bits + global_warp_id * tau; + + for(int i = 0; i < tau; i++) + { + // Step 10: if e[i] = 1 then t ← FMLM(t, y) + if(my_exp[i] == 1ULL) + { + XYfixWarpVector_dev( + t_buf, y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // Step 11: y ← FMLM(y, y) + XYfixWarpSquareVector_dev( + y_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // 去除 Montgomery 因子,还原真实结果 + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, + negmodn, negmodn_shoup, + modn, modn_shoup, + n_poly, mod, + twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, + NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Step 14: return t → 写回全局内存 + // ================================================================ + #pragma unroll + for(int j = 0; j < 8; j++) + output[off + 8*lane_id + j] = t_buf[8*lane_id + j]; +} +*/ + +__global__ void FMLMKernel_e( + const uint64_t *__restrict__ A, + const uint64_t + *__restrict__ exp_bits, // 格式改为压缩:[m * (tau/64)] uint64_t + int tau, uint64_t *output, int m, const uint64_t *__restrict__ d_r0, + uint64_t *d_r1_ct, uint64_t *d_r1_nct, uint64_t *negmodn, + uint64_t *negmodn_shoup, uint64_t *modn, uint64_t *modn_shoup, + uint64_t n_poly, uint64_t mod, const uint64_t *twiddle, + const uint64_t *twiddle_shoup, const uint64_t *ICTTwiddle, + const uint64_t *ICTTwiddle_shoup, const uint64_t *NCTtwiddle, + const uint64_t *NCTtwiddle_shoup, const uint64_t *InvNCTtwiddle, + const uint64_t *InvNCTtwiddle_shoup, uint64_t inv, uint64_t inv_shoup, + const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // ================================================================ + // 优化:压缩格式加载 + // + // exp_bits 格式:[m * (tau/64)] uint64_t + // 每个指数存为 tau/64 = 16 个 uint64_t(packed,LSB first) + // 第 p 个指数占 exp_bits[p*16 .. p*16+15] + // + // 加载方式: + // 每个 warp 只需加载 16 个 uint64_t = 128 字节(1 个 cache line) + // lane 0~15 各加载 1 个 uint64_t(合并访问) + // lane 16~31 加载 0(通过 __shfl_sync 获取) + // + // 总加载量:2048 warp × 128B = 256KB(完全驻留 L2 Cache) + // ================================================================ + const int exp_limbs = tau / 64; // = 16(tau=1024时) + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + // 每个 lane 持有一个 limb(lane 0持有limb0,lane 1持有limb1,...,lane + // 15持有limb15) lane 16~31 持有 0(不参与实际数据,通过 shfl 获取) + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) + my_limb = my_packed_exp[lane_id]; // 合并访问:lane 0~15 读连续地址 + +// ================================================================ +// Step 6: y ← x +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀ +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // Step 8: y ← FMLM(y, r₁) + // ================================================================ + XYfixWarpROneVector_dev(y_buf, d_r1_ct, d_r1_nct, negmodn, negmodn_shoup, + modn, modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + + // ================================================================ + // Steps 9–12: for i = 0 to τ-1 + // + // 取第 i 个 bit: + // limb 下标 = i / 64 → 该 limb 在 lane (i/64) 中 + // bit 位置 = i % 64 + // 通过 __shfl_sync 从 lane(i/64) 广播 my_limb 到所有 lane + // 取目标 bit,全程零全局内存访问 + // ================================================================ + for (int i = 0; i < tau; i++) { + // 从持有 limb(i/64) 的 lane 广播到全部 lane + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + + // ================================================================ + // Step 13: t ← FMLM(t, 1) + // ================================================================ + XYfixWarpIRVector_dev( + t_buf, negmodn, negmodn_shoup, modn, modn_shoup, n_poly, mod, twiddle, + twiddle_shoup, ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, inv_shoup, sample); + __syncwarp(); + +// ================================================================ +// Step 14: 写回全局内存 +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ── 参数 ──────────────────────────────────────────────────────── +#define TEST_BATCH 1000 // 测试用例数 +#define ARR_LEN 256 // 大数 limb 数 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 +#define TAU 1024 // 指数 bit 数 +#define EXP_U64_LIMBS (TAU / 64) // = 16,压缩格式每个指数占 16 个 uint64_t +#define N2_TOP_IDX 240 // n² 最高非零 limb 索引 +#define N2_TOP_VAL 27958ULL // n² 最高 limb 的值(上界) +#define WARP_PER_BLK 8 // 每 block 的 warp 数 +#define NUM_STREAMS 4 + +__global__ void FMLE_mod2_Kernel( + const uint64_t *__restrict__ A, // 输入:c̃₁(蒙哥马利域) + const uint64_t *__restrict__ exp_bits, // 指数 m(压缩格式) + int tau, // m 的 bit 长度 + uint64_t *output, // 输出:c₁^m · R(蒙哥马利域) + int m, // 测试用例数量 + const uint64_t *__restrict__ d_r0, // r₀ = R mod n²(1的蒙哥马利表示) + // ★ 去掉 d_r1_ct / d_r1_nct:Step8 已跳过,不再需要 + uint64_t *negmodn, uint64_t *negmodn_shoup, uint64_t *modn, + uint64_t *modn_shoup, uint64_t n_poly, uint64_t mod, + const uint64_t *twiddle, const uint64_t *twiddle_shoup, + const uint64_t *ICTTwiddle, const uint64_t *ICTTwiddle_shoup, + const uint64_t *NCTtwiddle, const uint64_t *NCTtwiddle_shoup, + const uint64_t *InvNCTtwiddle, const uint64_t *InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, const uint64_t *sample) { + const int warp_in_block = threadIdx.x >> 5; + const int warps_per_block = blockDim.x >> 5; + const int global_warp_id = blockIdx.x * warps_per_block + warp_in_block; + const uint32_t lane_id = threadIdx.x & 0x1F; + + if (global_warp_id >= m) return; + + extern __shared__ uint64_t smem[]; + uint64_t *const ws = smem + warp_in_block * 512; + uint64_t *const y_buf = ws; + uint64_t *const t_buf = ws + 256; + + const int off = global_warp_id * 256; + + // 压缩格式加载指数(与原版相同) + const int exp_limbs = tau / 64; + const uint64_t *my_packed_exp = exp_bits + (size_t)global_warp_id * exp_limbs; + + uint64_t my_limb = 0ULL; + if ((int)lane_id < exp_limbs) my_limb = my_packed_exp[lane_id]; + +// ================================================================ +// Step 6: y ← c̃₁(直接加载蒙哥马利域输入,与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) y_buf[8 * lane_id + j] = A[off + 8 * lane_id + j]; + __syncwarp(); + +// ================================================================ +// Step 7: t ← r₀(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) t_buf[8 * lane_id + j] = d_r0[8 * lane_id + j]; + __syncwarp(); + + // ================================================================ + // ★ Step 8:跳过! + // + // 原版:y ← XYfixWarpROneVector_dev(y, r₁) + // 即 y = y · R²· R⁻¹ = y · R,将普通域转为蒙哥马利域 + // + // FMLE_mod2:输入 A 已经是 c̃₁ = c₁ · R(蒙哥马利域) + // 无需转换,y_buf 已正确,直接进入循环 + // ================================================================ + + // ================================================================ + // Steps 9-12: 平方-乘循环(与原版完全相同) + // ================================================================ + for (int i = 0; i < tau; i++) { + uint64_t src_limb = __shfl_sync(0xFFFFFFFF, my_limb, (uint32_t)(i / 64)); + uint64_t the_bit = (src_limb >> (i % 64)) & 1ULL; + + if (the_bit == 1ULL) { + XYfixWarpVector_dev(t_buf, y_buf, negmodn, negmodn_shoup, modn, + modn_shoup, n_poly, mod, twiddle, twiddle_shoup, + ICTTwiddle, ICTTwiddle_shoup, NCTtwiddle, + NCTtwiddle_shoup, InvNCTtwiddle, InvNCTtwiddle_shoup, + inv, inv_shoup, sample); + __syncwarp(); + } + + XYfixWarpSquareVector_dev(y_buf, negmodn, negmodn_shoup, modn, modn_shoup, + n_poly, mod, twiddle, twiddle_shoup, ICTTwiddle, + ICTTwiddle_shoup, NCTtwiddle, NCTtwiddle_shoup, + InvNCTtwiddle, InvNCTtwiddle_shoup, inv, + inv_shoup, sample); + __syncwarp(); + } + +// ================================================================ +// ★ Step 13:跳过! +// +// 原版:t ← XYfixWarpIRVector_dev(t) +// 即 t = t · 1 · R⁻¹,将蒙哥马利域转回普通域 +// +// FMLE_mod2:保持蒙哥马利输出 t = c₁^m · R +// 供后续同态运算继续使用,无需转回普通域 +// ================================================================ + +// ================================================================ +// Step 14: 写回全局内存(与原版相同) +// ================================================================ +#pragma unroll + for (int j = 0; j < 8; j++) + output[off + 8 * lane_id + j] = t_buf[8 * lane_id + j]; +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct MulPlainTiming { + float kernel_ms; // 各流内核时间的最大值 + float pipeline_ms; // 流水线端到端时间(含流创建/销毁) +}; + +// ============================================================= +// paillier_mulplain +// +// 功能:批量 MulPlain(同态标量乘法),使用 CUDA 流并行 +// +// 要求: +// · 所有数据已在 GPU 上(设备指针),不做 H2D/D2H +// · 内部使用 NUM_STREAMS 条流,每流负责约 batch/NUM_STREAMS 个用例 +// · batch 无需是 NUM_STREAMS 的整数倍 +// · 计算过程:FMLE_mod2_Kernel(蒙哥马利域 → 蒙哥马利域) +// +// 流水线示意(NUM_STREAMS=4,batch 不整除时最后一流更少): +// Stream0: [Kernel, actual_chunk[0] 个] +// Stream1: [Kernel, actual_chunk[1] 个] +// Stream2: [Kernel, actual_chunk[2] 个] +// Stream3: [Kernel, actual_chunk[3] 个] ← 可能比其他流少,甚至为 0(跳过) +// +// 参数: +// d_ciphertexts [batch * ARR_LEN] 蒙哥马利域密文输入(设备) +// d_exps [batch * EXP_U64_LIMBS] 压缩指数(设备) +// d_output [batch * ARR_LEN] 蒙哥马利域输出(设备) +// 其余参数为预计算常量(已在 GPU 上) +// ============================================================= +MulPlainTiming paillier_mulplain( + const uint64_t *d_ciphertexts, const uint64_t *d_exps, uint64_t *d_output, + int batch, int tau, const uint64_t *d_r0, uint64_t *d_negmodn, + uint64_t *d_negmodn_shoup, uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, const uint64_t *d_twiddle, const uint64_t *d_twiddle_shoup, + const uint64_t *d_ICTTwiddle, const uint64_t *d_ICTTwiddle_shoup, + const uint64_t *d_NCTtwiddle, const uint64_t *d_NCTtwiddle_shoup, + const uint64_t *d_InvNCTtwiddle, const uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, const uint64_t *d_sample) { + // ── 向上取整 chunk,无需 batch 整除 NUM_STREAMS ────────────── + // 改动1:向上取整,保证所有用例被覆盖 + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + // ── 创建流和事件 ────────────────────────────────────────────── + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev_ker[NUM_STREAMS][2]; + + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][0])); + CUDA_CHECK(cudaEventCreate(&ev_ker[s][1])); + } + + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + // ── 流水线提交 ──────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + // 改动2:计算本流起点,超出 batch 则无任务,跳过 + int offset = s * chunk; + if (offset >= batch) break; + + // 改动3:本流实际处理量(最后一流可能不足 chunk) + int actual_chunk = ((offset + chunk) > batch) ? (batch - offset) : chunk; + + // 改动4:按 actual_chunk 计算 blocks,不会启动多余 warp + const int blocks = (actual_chunk + WARP_PER_BLK - 1) / WARP_PER_BLK; + + // 偏移(元素数) + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * EXP_U64_LIMBS; + + CUDA_CHECK(cudaEventRecord(ev_ker[s][0], streams[s])); + + // 改动5:传入 actual_chunk 而非 chunk,内核 guard 不会越界 + FMLE_mod2_Kernel<<>>( + d_ciphertexts + ct_off, d_exps + exp_off, tau, d_output + ct_off, + actual_chunk, // ← 实际用例数 + d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + + CUDA_CHECK(cudaEventRecord(ev_ker[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + // ── 计算各流内核时间,取最大值 ─────────────────────────────── + MulPlainTiming t = {0.f, 0.f}; + for (int s = 0; s < NUM_STREAMS; s++) { + // 未启动的流(offset >= batch 时 break 前已跳过)事件未记录, + // 只统计实际启动的流 + int offset = s * chunk; + if (offset >= batch) break; + + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev_ker[s][0], ev_ker[s][1])); + if (ms > t.kernel_ms) t.kernel_ms = ms; + } + CUDA_CHECK(cudaEventElapsedTime(&t.pipeline_ms, ev_pipe_start, ev_pipe_end)); + + // ── 清理流和事件 ────────────────────────────────────────────── + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev_ker[s][0]); + cudaEventDestroy(ev_ker[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + + return t; +} + +// ============================================================= +// 写结果到 txt(与原版格式相同) +// ============================================================= +static void write_to_txt(const char *filename, int num, const uint64_t *h_bases, + const uint64_t *h_exps, const uint64_t *h_results, + const uint64_t *modn) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d %d\n", num, ARR_LEN, BASE_BITS, TAU); + + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int i = 0; i < num; i++) { + const uint64_t *x = h_bases + (size_t)i * ARR_LEN; + const uint64_t *e = h_exps + (size_t)i * EXP_U64_LIMBS; + const uint64_t *r = h_results + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)x[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < EXP_U64_LIMBS; j++) { + fprintf(fp, "%llu", (unsigned long long)e[j]); + fprintf(fp, j < EXP_U64_LIMBS - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)r[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ================================================================ +// paillier_cpowm +// +// 并行计算 batch 对 c_i^m_i mod n^2 +// +// 输入: +// h_ciphertexts [batch * 256] 密文,每个大数 256 个 17-bit limbs +// h_plaintexts [batch * 16] 明文,每个大数 16 个 64-bit +// limbs(共1024bit) batch 明密文对数量 +// 以下为预计算好的 GPU 参数(均已在 device 端) +// 输出: +// h_output [batch * 256] c^m mod n^2 结果 +// ================================================================ +void paillier_cpowm(const uint64_t *h_ciphertexts, const uint64_t *h_plaintexts, + int batch, uint64_t *d_r0, uint64_t *d_r1_ct, + uint64_t *d_r1_nct, uint64_t *d_negModn, + uint64_t *d_negModn_shoup, uint64_t *d_Modn, + uint64_t *d_Modn_shoup, uint64_t n_poly, uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle, uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv, uint64_t inv_shoup, uint64_t *d_sample, + uint64_t *h_output) { + const int cpowm_tau = 1024; + const int cpowm_exp_u64_limbs = 16; + const int cpowm_limbs_per_num = 256; + const int cpowm_warp_per_blk = 4; + const int cpowm_threads_per_blk = cpowm_warp_per_blk * 32; + + const size_t cpowm_ct_bytes = + (size_t)batch * cpowm_limbs_per_num * sizeof(uint64_t); + const size_t cpowm_expbits_bytes = + (size_t)batch * cpowm_tau * sizeof(uint64_t); + + // ---- Step 1:明文展开为 LSB-first bit 数组 ---- + uint64_t *exp_bits_host = (uint64_t *)malloc(cpowm_expbits_bytes); + if (!exp_bits_host) { + fprintf(stderr, "[paillier_cpowm] malloc exp_bits_host 失败\n"); + return; + } + + for (int p = 0; p < batch; p++) { + const uint64_t *e = h_plaintexts + (size_t)p * cpowm_exp_u64_limbs; + uint64_t *bits = exp_bits_host + (size_t)p * cpowm_tau; + for (int i = 0; i < cpowm_tau; i++) { + int limb_idx = i / 64; + int bit_idx = i % 64; + bits[i] = (e[limb_idx] >> bit_idx) & 1ULL; + } + } + + // ---- Step 2:分配 device 内存并传输 ---- + uint64_t *d_ciphertexts = NULL; + uint64_t *d_output = NULL; + uint64_t *d_exp_bits = NULL; + + cudaMalloc(&d_ciphertexts, cpowm_ct_bytes); + cudaMalloc(&d_output, cpowm_ct_bytes); + cudaMalloc(&d_exp_bits, cpowm_expbits_bytes); + + cudaMemcpy(d_ciphertexts, h_ciphertexts, cpowm_ct_bytes, + cudaMemcpyHostToDevice); + cudaMemcpy(d_exp_bits, exp_bits_host, cpowm_expbits_bytes, + cudaMemcpyHostToDevice); + + cudaError_t err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内存传输失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + + // ---- Step 3:调用内核 ---- + { + int n_blocks = (batch + cpowm_warp_per_blk - 1) / cpowm_warp_per_blk; + size_t smem_size = (size_t)cpowm_warp_per_blk * 512 * sizeof(uint64_t); + + FMLMKernel_e<<>>( + d_ciphertexts, d_exp_bits, cpowm_tau, d_output, batch, d_r0, d_r1_ct, + d_r1_nct, d_negModn, d_negModn_shoup, d_Modn, d_Modn_shoup, n_poly, mod, + d_twiddle, d_twiddle_shoup, d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv, inv_shoup, d_sample); + + cudaDeviceSynchronize(); + err = cudaGetLastError(); + if (err != cudaSuccess) { + fprintf(stderr, "[paillier_cpowm] 内核执行失败: %s\n", + cudaGetErrorString(err)); + goto cpowm_cleanup; + } + } + + // ---- Step 4:拷回结果 ---- + cudaMemcpy(h_output, d_output, cpowm_ct_bytes, cudaMemcpyDeviceToHost); + +cpowm_cleanup: + free(exp_bits_host); + if (d_ciphertexts) cudaFree(d_ciphertexts); + if (d_output) cudaFree(d_output); + if (d_exp_bits) cudaFree(d_exp_bits); +} + +#define NUM_TESTS 2000000 +#define ARR_LEN 256 +#define BASE_BITS 17 +#define BASE (1ULL << BASE_BITS) // 131072 + +// n 是 2048-bit,在 base 2^17 下约占 121 个 limb +// 超出部分为 0,无需额外处理(循环自然跳过) +#define N_WORDS 121 + +#ifndef THREADS_PER_BLOCK +#define THREADS_PER_BLOCK (WARPS_PER_BLOCK * 32) +#endif + +#define NUM_STREAMS 4 +#define CHUNK (NUM_TESTS / NUM_STREAMS) + +static void gen_random_bignum(uint64_t *out, int max_nonzero) { + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_nonzero) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ------------------------------------------------------------- +// 随机大数生成(limb 限制在 n² 有效范围内) +// modn:n² 的 limb 表示,确定有效 limb 上界 +// ------------------------------------------------------------- +static void generate_random_bignum(uint64_t *out, const uint64_t *modn) { + // 找 modn 最高非零 limb 的位置 + int max_idx = 0; + for (int i = ARR_LEN - 1; i >= 0; i--) { + if (modn[i] != 0) { + max_idx = i; + break; + } + } + // 只填充低于 max_idx 的 limb,保证生成值 < n² + for (int i = 0; i < ARR_LEN; i++) { + if (i < max_idx) { + uint64_t lo = (uint64_t)rand(); + uint64_t hi = (uint64_t)rand(); + out[i] = ((hi << 15) | lo) % BASE; + } else { + out[i] = 0; + } + } +} + +// ============================================================= +// GPU 端 Step1:批量计算 (1 ± m*n) mod n² +// +// 线程布局: +// gridDim.x = NUM_TESTS(每个 block 负责一个测试用例) +// blockDim.x = ARR_LEN (每个 block 的线程数 = limb 数量) +// +// 并行策略: +// ├─ 并行阶段:线程 i 计算 prod[i] = Σ m[j]·n[i-j](不含进位) +// │ 每项 m[j]·n[k] < 2^34,累加 256 项 < 2^42,不溢出 uint64_t +// └─ 串行阶段:线程 0 做进位传播 + ±1 调整 +// (各 block 间独立,不影响总并行度) +// ============================================================= +__global__ void compute_plain_factor_batch_kernel( + const uint64_t *__restrict__ d_m_batch, // [NUM_TESTS * ARR_LEN] 明文批 + const uint64_t *__restrict__ d_n, // [ARR_LEN] 公钥 n + const uint64_t *__restrict__ d_n2, // [ARR_LEN] 模数 n² + uint64_t *d_factor, // [NUM_TESTS * ARR_LEN] 输出 + int sign // +1: AddPlain, -1: SubPlain +) { + int test_id = blockIdx.x; + int i = threadIdx.x; // 当前线程负责第 i 个 limb + + if (test_id >= NUM_TESTS || i >= ARR_LEN) return; + + const uint64_t *m = d_m_batch + (size_t)test_id * ARR_LEN; + uint64_t *out = d_factor + (size_t)test_id * ARR_LEN; + + // ── 共享内存:m*n 的乘积(不含进位的中间值)────────────────── + __shared__ uint64_t prod[ARR_LEN]; + + // ── 并行阶段:计算第 i 个 limb 的偏积和 ────────────────────── + // prod[i] = Σ_{j=0}^{i} m[j] * n[i-j] + // (m[j]=0 for j≥N_WORDS, n[k]=0 for k≥N_WORDS,自动截断) + uint64_t sum = 0; + for (int j = 0; j <= i; j++) { + // j < N_WORDS 且 (i-j) < N_WORDS 才有非零贡献 + sum += m[j] * d_n[i - j]; + } + prod[i] = sum; + + __syncthreads(); + + // ── 串行阶段(仅线程 0):进位传播 + ±1 ───────────────────── + if (i == 0) { + // ① 进位传播(m*n 的 limb 规格化到 [0, BASE)) + for (int k = 0; k < ARR_LEN - 1; k++) { + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + if (sign > 0) { + // ② AddPlain:prod = mn,加 1 → (1+mn) + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + + } else { + // ② SubPlain:(1-mn) mod n² = n² - mn + 1(因 mn > 1) + // Step A: n² - mn(大数减法) + int64_t borrow = 0; + for (int k = 0; k < ARR_LEN; k++) { + int64_t diff = (int64_t)d_n2[k] - (int64_t)prod[k] - borrow; + if (diff < 0) { + prod[k] = (uint64_t)(diff + (int64_t)BASE); + borrow = 1; + } else { + prod[k] = (uint64_t)diff; + borrow = 0; + } + } + // Step B: + 1 + prod[0] += 1; + for (int k = 0; k < ARR_LEN - 1; k++) { + if (prod[k] < BASE) break; + prod[k + 1] += prod[k] >> BASE_BITS; + prod[k] &= BASE - 1; + } + } + + // ③ 写出结果 + for (int k = 0; k < ARR_LEN; k++) out[k] = prod[k]; + } +} + +// ============================================================= +// 计时结构体 +// ============================================================= +struct StepTime { + float step1_ms; // compute_plain_factor (GPU bignum 乘法) + float step2_ms; // XYfixWarpROneVector(转入蒙哥马利域) + float step3_ms; // XYfixWarpVector(与 c̃1 相乘) + float h2d_ms; // 数据上传 + float d2h_ms; // 结果下载 +}; + +/* +// ============================================================= +// paillier_addplain_subplain_batch +// 完整批量运算:H2D → Step1 → Step2 → Step3 → D2H +// ============================================================= +StepTime paillier_addplain_subplain_batch( + uint64_t *h_c1_tilde, // [NUM_TESTS*ARR_LEN] 蒙哥马利密文输入 (pinned) + uint64_t *h_m_batch, // [NUM_TESTS*ARR_LEN] 明文 m 批 (pinned) + uint64_t *h_result, // [NUM_TESTS*ARR_LEN] 输出 (pinned) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, + uint64_t *d_modn, uint64_t *d_modn_shoup, + uint64_t mod, + uint64_t *d_twiddle, uint64_t *d_twiddle_shoup, + uint64_t *d_ICTTwiddle, uint64_t *d_ICTTwiddle_shoup, + uint64_t *d_NCTtwiddle, uint64_t *d_NCTtwiddle_shoup, + uint64_t *d_InvNCTtwiddle,uint64_t *d_InvNCTtwiddle_shoup, + uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) +{ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + StepTime t = {0,0,0,0,0}; + + // ── 分配 GPU 工作缓冲区 ──────────────────────────────────────── + // d_factor_mont 不再需要:XYfixWarpROneVector 原地写回 d_factor + // d_R2_batch 不再需要:R² 以 d_r1_ct/d_r1_nct 直接传入 + uint64_t *d_c1, *d_m, *d_factor; + CUDA_CHECK(cudaMalloc(&d_c1, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); // (1±mn) → 原地转换为 +(1̃±mn) + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // H2D:上传 c̃1 和 m 批 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(d_c1, h_c1_tilde, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaMemcpy(d_m, h_m_batch, batch_bytes, +cudaMemcpyHostToDevice)); CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.h2d_ms, e0, e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // + // 每个 block = 一个测试用例 + // 线程 i 并行计算 prod[i],线程 0 做进位和±1调整 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>( + d_m, d_n, d_n2, d_factor, sign + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:(1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector(factor, d_r1_ct, d_r1_nct) + // 等价于 FMLM((1±mn), R²) = (1±mn)·R + // + // R² 以预存的 CT/NCT 形式直接传入,无需 broadcast + // 结果原地写回 d_factor:普通域 → 蒙哥马利域 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, // inout:(1±mn) → (1̃±mn),原地写回 + d_r1_ct, // inoutAct:R² 的 CT 形式 + d_r1_nct, // inoutAnct:R² 的 NCT 形式 + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:c̃1 与 (1̃±mn) 相乘 + // FMLM( c̃1, (1̃±mn) ) = c1·R · (1±mn)·R · R⁻¹ + // = c1·(1±mn)·R ✓ + // + // d_factor 此时已是蒙哥马利形式(Step2 原地写回) + // d_c1 存放 c̃1,结果写回 d_c1 + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, // FMLM(c̃1, (1̃±mn)) + d_negmodn, d_negmodn_shoup, + d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, + d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, + d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, + inv_val, inv_shoup_val, d_sample + ); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + // ============================================================ + // D2H:结果回传(d_c1 已被 Step3 写入结果) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + CUDA_CHECK(cudaMemcpy(h_result, d_c1, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.d2h_ms, e0, e1)); + + cudaEventDestroy(e0); cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_c1)); + CUDA_CHECK(cudaFree(d_m)); + CUDA_CHECK(cudaFree(d_factor)); + + return t; +} +*/ + +// ============================================================ +// 修改后的 paillier_addplain_subplain_batch +// +// 变化: +// · 不再接收 h_c1_tilde / h_m_batch / h_result(主机指针) +// · 改为接收 d_c1 / d_m(已在 GPU 上的设备指针) +// · 内部不做 H2D / D2H,也不 cudaMalloc/cudaFree d_c1 和 d_m +// · d_c1 同时作为输入(c̃1)和输出(Step3 结果写回) +// · d_factor 仍在函数内部分配/释放(临时工作缓冲) +// · h2d_ms / d2h_ms 固定为 0,由外部计时 +// ============================================================ +StepTime paillier_addplain_subplain_batch( + uint64_t *d_c1, // [NUM_TESTS*ARR_LEN] device: 蒙哥马利密文(输入兼输出) + uint64_t *d_m, // [NUM_TESTS*ARR_LEN] device: 明文 m 批(已上传) + uint64_t *d_n, // [ARR_LEN] 公钥 n(已在 GPU) + uint64_t *d_n2, // [ARR_LEN] 模数 n²(已在 GPU) + uint64_t *d_r1_ct, // [ARR_LEN] R² mod n² 的 CT 形式(已在 GPU) + uint64_t *d_r1_nct, // [ARR_LEN] R² mod n² 的 NCT 形式(已在 GPU) + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, uint64_t *d_twiddle, + uint64_t *d_twiddle_shoup, uint64_t *d_ICTTwiddle, + uint64_t *d_ICTTwiddle_shoup, uint64_t *d_NCTtwiddle, + uint64_t *d_NCTtwiddle_shoup, uint64_t *d_InvNCTtwiddle, + uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, uint64_t inv_shoup_val, + uint64_t *d_sample, + int sign // +1: AddPlain, -1: SubPlain +) { + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + const int blocks_per_batch = + (NUM_TESTS + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + dim3 fmlm_grid(blocks_per_batch); + dim3 fmlm_block(THREADS_PER_BLOCK); + + // h2d_ms / d2h_ms 由外部测量,此处置 0 + StepTime t = {0, 0, 0, 0, 0}; + + // 只分配内部临时工作缓冲 d_factor((1±mn) 及其蒙哥马利形式) + // d_c1、d_m 由调用者传入,不在此处分配或释放 + uint64_t *d_factor; + CUDA_CHECK(cudaMalloc(&d_factor, batch_bytes)); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ============================================================ + // Step1:计算 (1 ± m*n) mod n²(GPU 大数乘法) + // 输入:d_m(明文批)、d_n(公钥 n)、d_n2(模数 n²) + // 输出:d_factor = (1 ± m*n) mod n² + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + compute_plain_factor_batch_kernel<<>>(d_m, d_n, d_n2, + d_factor, sign); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ============================================================ + // Step2:将 (1±mn) 转入蒙哥马利域 + // XYfixWarpROneVector 等价于 FMLM((1±mn), R²) = (1±mn)·R + // 结果原地写回 d_factor + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<>>( + d_factor, d_r1_ct, d_r1_nct, d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + + // ============================================================ + // Step3:FMLM(d_c1, d_factor) = c1*(1±mn)*R + // 结果写回 d_c1(同一缓冲区,调用者读取 D2H 后即为最终输出) + // ============================================================ + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_factor, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step3_ms, e0, e1)); + + cudaEventDestroy(e0); + cudaEventDestroy(e1); + CUDA_CHECK(cudaFree(d_factor)); + // d_c1 和 d_m 由调用者负责 cudaFree + + return t; +} + +// ============================================================ +// 打印计时报告(含外部测量的 H2D / D2H) +// ============================================================ +static void print_steptime_ext(const char *label, const StepTime &t, + float h2d_ms, // 外部测量的 H2D 时间 + float d2h_ms // 外部测量的 D2H 时间 +) { + float compute_ms = t.step1_ms + t.step2_ms + t.step3_ms; + float total_ms = h2d_ms + compute_ms + d2h_ms; + + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" [外部] H2D 上传 : %8.4f ms\n", h2d_ms); + printf(" Step1 (1+-mn) GPU 大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector]\n", + t.step3_ms); + printf(" [外部] D2H 下载 : %8.4f ms\n", d2h_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total_ms); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + compute_ms * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total_ms * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 打印计时报告 +// ============================================================= +static void print_steptime(const char *label, const StepTime &t) { + float total = t.h2d_ms + t.step1_ms + t.step2_ms + t.step3_ms + t.d2h_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" H2D 上传 : %8.4f ms\n", t.h2d_ms); + printf(" Step1 (1±mn) GPU大数乘 : %8.4f ms\n", t.step1_ms); + printf(" Step2 转蒙哥马利域 : %8.4f ms [XYfixWarpROneVector]\n", + t.step2_ms); + printf(" Step3 与密文相乘 : %8.4f ms [XYfixWarpVector(c̃1, factor)]\n", + t.step3_ms); + printf(" D2H 下载 : %8.4f ms\n", t.d2h_ms); + printf(" ──────────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 Step1+2+3 : %8.4f us/次\n", + (t.step1_ms + t.step2_ms + t.step3_ms) * 1000.f / NUM_TESTS); + printf(" 平均单次端到端 : %8.4f us/次\n", + total * 1000.f / NUM_TESTS); + printf("======================================================\n"); +} + +// ============================================================= +// 写 txt 供 Python 验证 +// 格式: +// 行1:NUM_TESTS ARR_LEN BASE_BITS +// 行2:Modn (n²) limbs +// 行3:N_arr (n) limbs ← 验证需要 n +// 每个测试用例 3 行:m / c1_tilde / result +// ============================================================= +static void write_results_to_txt(const char *filename, + const uint64_t *h_m_batch, + const uint64_t *h_c1_tilde, + const uint64_t *h_result, const uint64_t *modn, + const uint64_t *n_arr) { + FILE *fp = fopen(filename, "w"); + if (!fp) { + fprintf(stderr, "无法创建 %s\n", filename); + return; + } + + fprintf(fp, "%d %d %d\n", NUM_TESTS, ARR_LEN, BASE_BITS); + + // n² limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)modn[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + // n limbs + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)n_arr[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + + for (int i = 0; i < NUM_TESTS; i++) { + const uint64_t *m = h_m_batch + (size_t)i * ARR_LEN; + const uint64_t *c1 = h_c1_tilde + (size_t)i * ARR_LEN; + const uint64_t *res = h_result + (size_t)i * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)m[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)c1[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + for (int j = 0; j < ARR_LEN; j++) { + fprintf(fp, "%llu", (unsigned long long)res[j]); + fprintf(fp, j < ARR_LEN - 1 ? " " : "\n"); + } + } + fclose(fp); + printf("结果已写入 %s\n", filename); +} + +// ============================================================= +// broadcast_fill_kernel +// 将单个 ARR_LEN 的大数广播到 batch × ARR_LEN 的目标批量数组 +// dst[test_id * ARR_LEN + limb] = src[limb] +// ============================================================= +__global__ void broadcast_fill_kernel(uint64_t *dst, const uint64_t *src, + int batch) { + int tid = blockIdx.x * blockDim.x + threadIdx.x; + if (tid >= batch * ARR_LEN) return; + dst[tid] = src[tid % ARR_LEN]; +} + +// ============================================================= +// paillier_randomize 计时结构体 +// ============================================================= +struct RandomizeTiming { + float step1_ms; // h_s → h̃_s(XYfixWarpROneVector,1 warp) + float step2_ms; // 广播(broadcast_fill_kernel) + float step3_ms; // 批量模幂(FMLE_mod2_Kernel,流水线最大值) + float step4_ms; // 批量乘法(XYfixWarpVector) +}; + +// ============================================================= +// paillier_randomize +// +// 功能:批量 Randomize 同态操作 +// c̃_out[i] = c₁[i] · h_s^{r_i} · R mod n² (蒙哥马利形式) +// +// 数据流: +// Step① XYfixWarpROneVector: h_s → h̃_s = h_s · R (1 warp) +// Step② broadcast_fill_kernel: h̃_s → d_hs_batch[batch] +// Step③ FMLE_mod2_Kernel: h̃_s^{r_i}(蒙→蒙,多流) +// Step④ XYfixWarpVector: c̃₁[i] · h̃_s^{r_i}(结果→c̃_out) +// +// 注意: +// · 所有指针均为设备指针,函数内部不做 H2D / D2H +// · d_hs_single 会被原地覆写为 h̃_s = h_s · R(调用者如需保留原值需自行备份) +// · d_c1 既是输入也是输出(原地更新) +// · d_r_batch 格式:batch × (tau/64) 个 uint64_t,压缩存储每个 r_i +// ============================================================= +RandomizeTiming paillier_randomize( + uint64_t *d_c1, // [batch × ARR_LEN] 蒙哥马利密文(输入兼输出) + uint64_t *d_hs_single, // [ARR_LEN] h_s 普通域(原地改写为 h̃_s) + const uint64_t *d_r_batch, // [batch × (tau/64)] 压缩随机指数 r_i + int batch, + int tau, // 指数 bit 长度 + const uint64_t *d_r0, // [ARR_LEN] R mod n²(1 的蒙哥马利形式) + const uint64_t *d_r1_ct, // [ARR_LEN] R² 的 CT 形式 + const uint64_t *d_r1_nct, // [ARR_LEN] R² 的 NCT 形式 + uint64_t *d_negmodn, uint64_t *d_negmodn_shoup, uint64_t *d_modn, + uint64_t *d_modn_shoup, uint64_t mod, const uint64_t *d_twiddle, + const uint64_t *d_twiddle_shoup, const uint64_t *d_ICTTwiddle, + const uint64_t *d_ICTTwiddle_shoup, const uint64_t *d_NCTtwiddle, + const uint64_t *d_NCTtwiddle_shoup, const uint64_t *d_InvNCTtwiddle, + const uint64_t *d_InvNCTtwiddle_shoup, uint64_t inv_val, + uint64_t inv_shoup_val, const uint64_t *d_sample) { + RandomizeTiming t = {0.f, 0.f, 0.f, 0.f}; + const size_t batch_bytes = (size_t)batch * ARR_LEN * sizeof(uint64_t); + + cudaEvent_t e0, e1; + CUDA_CHECK(cudaEventCreate(&e0)); + CUDA_CHECK(cudaEventCreate(&e1)); + + // ───────────────────────────────────────────────────────────── + // Step①:h_s → h̃_s = h_s · R mod n² + // 1 block × 32 threads(单 warp)覆盖 ARR_LEN=256 个 limb + // 结果原地写回 d_hs_single + // ───────────────────────────────────────────────────────────── + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpROneVector<<<1, 32>>>( + d_hs_single, const_cast(d_r1_ct), + const_cast(d_r1_nct), d_negmodn, d_negmodn_shoup, d_modn, + d_modn_shoup, (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, + d_ICTTwiddle, d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, + d_InvNCTtwiddle, d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step1_ms, e0, e1)); + + // ───────────────────────────────────────────────────────────── + // Step②:广播 h̃_s[ARR_LEN] → d_hs_batch[batch × ARR_LEN] + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_batch; + CUDA_CHECK(cudaMalloc(&d_hs_batch, batch_bytes)); + { + const int total = batch * ARR_LEN; + const int blk = 256; + const int grd = (total + blk - 1) / blk; + CUDA_CHECK(cudaEventRecord(e0)); + broadcast_fill_kernel<<>>(d_hs_batch, d_hs_single, batch); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step2_ms, e0, e1)); + } + + // ───────────────────────────────────────────────────────────── + // Step③:批量模幂 h̃_s^{r_i}(FMLE_mod2,蒙哥马利→蒙哥马利) + // 使用 NUM_STREAMS 条流并发,每流处理约 batch/NUM_STREAMS 个用例 + // ───────────────────────────────────────────────────────────── + uint64_t *d_hs_r; + CUDA_CHECK(cudaMalloc(&d_hs_r, batch_bytes)); + { + const int chunk = (batch + NUM_STREAMS - 1) / NUM_STREAMS; + const int exp_limbs = tau / 64; + const size_t smem_size = (size_t)WARP_PER_BLK * 512 * sizeof(uint64_t); + + cudaStream_t streams[NUM_STREAMS]; + cudaEvent_t ev[NUM_STREAMS][2]; + for (int s = 0; s < NUM_STREAMS; s++) { + CUDA_CHECK(cudaStreamCreate(&streams[s])); + CUDA_CHECK(cudaEventCreate(&ev[s][0])); + CUDA_CHECK(cudaEventCreate(&ev[s][1])); + } + cudaEvent_t ev_pipe_start, ev_pipe_end; + CUDA_CHECK(cudaEventCreate(&ev_pipe_start)); + CUDA_CHECK(cudaEventCreate(&ev_pipe_end)); + + CUDA_CHECK(cudaEventRecord(ev_pipe_start, 0)); + + for (int s = 0; s < NUM_STREAMS; s++) { + int offset = s * chunk; + if (offset >= batch) break; + int actual = (offset + chunk > batch) ? (batch - offset) : chunk; + int blocks = (actual + WARP_PER_BLK - 1) / WARP_PER_BLK; + + size_t ct_off = (size_t)offset * ARR_LEN; + size_t exp_off = (size_t)offset * exp_limbs; + + CUDA_CHECK(cudaEventRecord(ev[s][0], streams[s])); + FMLE_mod2_Kernel<<>>( + d_hs_batch + ct_off, d_r_batch + exp_off, tau, d_hs_r + ct_off, + actual, d_r0, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaEventRecord(ev[s][1], streams[s])); + } + + CUDA_CHECK(cudaEventRecord(ev_pipe_end, 0)); + CUDA_CHECK(cudaEventSynchronize(ev_pipe_end)); + CUDA_CHECK(cudaGetLastError()); + + float max_ms = 0.f; + for (int s = 0; s < NUM_STREAMS; s++) { + if (s * chunk >= batch) break; + float ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&ms, ev[s][0], ev[s][1])); + if (ms > max_ms) max_ms = ms; + } + t.step3_ms = max_ms; + + for (int s = 0; s < NUM_STREAMS; s++) { + cudaStreamDestroy(streams[s]); + cudaEventDestroy(ev[s][0]); + cudaEventDestroy(ev[s][1]); + } + cudaEventDestroy(ev_pipe_start); + cudaEventDestroy(ev_pipe_end); + } + CUDA_CHECK(cudaFree(d_hs_batch)); + + // ───────────────────────────────────────────────────────────── + // Step④:c̃_out[i] = FMLM(c̃₁[i], h̃_s^{r_i}) = c₁[i] · h_s^{r_i} · R + // 结果原地写回 d_c1 + // ───────────────────────────────────────────────────────────── + { + const int blocks = (batch + WARPS_PER_BLOCK - 1) / WARPS_PER_BLOCK; + CUDA_CHECK(cudaEventRecord(e0)); + XYfixWarpVector<<>>( + d_c1, d_hs_r, d_negmodn, d_negmodn_shoup, d_modn, d_modn_shoup, + (uint64_t)ARR_LEN, mod, d_twiddle, d_twiddle_shoup, d_ICTTwiddle, + d_ICTTwiddle_shoup, d_NCTtwiddle, d_NCTtwiddle_shoup, d_InvNCTtwiddle, + d_InvNCTtwiddle_shoup, inv_val, inv_shoup_val, d_sample); + CUDA_CHECK(cudaGetLastError()); + CUDA_CHECK(cudaEventRecord(e1)); + CUDA_CHECK(cudaEventSynchronize(e1)); + CUDA_CHECK(cudaEventElapsedTime(&t.step4_ms, e0, e1)); + } + + CUDA_CHECK(cudaFree(d_hs_r)); + CUDA_CHECK(cudaEventDestroy(e0)); + CUDA_CHECK(cudaEventDestroy(e1)); + + return t; +} + +// ============================================================= +// 打印 Randomize 计时报告 +// ============================================================= +static void print_randomize_timing(const char *label, const RandomizeTiming &t, + int batch) { + float total = t.step1_ms + t.step2_ms + t.step3_ms + t.step4_ms; + printf("\n======================================================\n"); + printf(" %s 计时报告\n", label); + printf("======================================================\n"); + printf(" Step1 h_s → h̃_s(1 warp) : %8.4f ms\n", t.step1_ms); + printf(" Step2 广播 h̃_s : %8.4f ms\n", t.step2_ms); + printf(" Step3 批量模幂 h̃_s^r : %8.4f ms [FMLE_mod2_Kernel]\n", + t.step3_ms); + printf(" Step4 批量乘法 c̃₁ · h̃_s^r : %8.4f ms [XYfixWarpVector]\n", + t.step4_ms); + printf(" ──────────────────────────────────────────\n"); + printf(" 端到端总时间 : %8.4f ms\n", total); + printf(" 平均单次 : %8.4f us/次\n", + total * 1000.f / batch); + printf("======================================================\n"); +} + +int main() { + uint64_t Modn[256] = { + 2649, 107256, 2340, 7668, 50437, 89753, 117294, 34082, 83769, + 106490, 49932, 81014, 10796, 34292, 88272, 35871, 1224, 70218, + 72484, 90481, 9348, 59016, 20888, 61690, 124, 21004, 97475, + 905, 112940, 56833, 106121, 126381, 83117, 119497, 72617, 85847, + 112617, 68151, 38906, 110574, 6475, 31434, 5205, 105833, 121951, + 28791, 57007, 122812, 74953, 23473, 98222, 45790, 91691, 117397, + 110559, 19813, 109479, 123726, 1814, 68875, 26386, 97607, 4246, + 117792, 20567, 14111, 99272, 120821, 70606, 27709, 106158, 56805, + 36973, 115639, 18908, 61130, 43630, 63374, 22728, 4738, 91891, + 97888, 98463, 6143, 26241, 16864, 103392, 28734, 112371, 25077, + 4971, 44003, 95710, 127770, 116422, 86796, 53898, 106151, 110768, + 4172, 16754, 115838, 102520, 23673, 31694, 9309, 70081, 4025, + 12143, 74312, 9125, 48862, 92677, 54421, 6506, 70890, 65534, + 103844, 63591, 71250, 42737, 107501, 69713, 128716, 73161, 72416, + 23904, 89050, 125203, 82817, 127482, 17228, 58669, 21454, 52151, + 18485, 51318, 121322, 104184, 54791, 18882, 119526, 74931, 90359, + 87105, 107810, 48763, 41791, 28610, 125248, 79695, 69759, 26996, + 68191, 57603, 66348, 105108, 91481, 80701, 38092, 51920, 77610, + 91327, 57269, 2646, 28099, 85966, 16858, 63989, 76676, 10478, + 9393, 102553, 62955, 48269, 127262, 80638, 109877, 86149, 80042, + 76558, 23669, 93286, 63134, 69813, 101873, 117146, 20027, 118499, + 123171, 125741, 42197, 23070, 36105, 39520, 70168, 58742, 75228, + 111900, 58941, 20733, 3481, 86208, 112571, 94728, 122773, 103601, + 112257, 89390, 108130, 1539, 74379, 115661, 126079, 27686, 31430, + 66734, 57259, 14251, 4473, 82338, 81921, 19921, 105501, 45514, + 68724, 69497, 8615, 98992, 113627, 81449, 55998, 102347, 25801, + 20790, 101333, 67616, 42614, 130518, 19795, 27958, 0, 0, + 0, 0, 0, 0, 0, 0, 0, 0, 0, + 0, 0, 0, 0}; + + uint64_t negModn[256] = { + 99547, 122922, 54553, 19080, 126680, 105037, 100182, 32356, 92056, + 117188, 107847, 59449, 108351, 48087, 82800, 122633, 83740, 58351, + 88003, 30380, 71079, 506, 81403, 131028, 36809, 103471, 120931, + 91211, 3573, 28302, 36730, 109307, 65780, 94975, 51906, 7151, + 122366, 93383, 130847, 82394, 91081, 86111, 23836, 16026, 38884, + 90913, 64711, 45361, 93021, 47087, 58737, 47333, 114937, 101384, + 101884, 49344, 63760, 42692, 53665, 110133, 35501, 63140, 109709, + 89509, 36084, 54908, 51440, 117163, 124796, 1277, 29552, 10841, + 61416, 43106, 50143, 96335, 25765, 102230, 30619, 97194, 36585, + 106032, 923, 62621, 56847, 93152, 44694, 35555, 51373, 60882, + 85910, 68903, 5365, 28331, 55843, 39632, 87996, 74917, 22652, + 74639, 105498, 7496, 36812, 87767, 122208, 16698, 90622, 59916, + 50048, 19896, 7951, 59415, 82759, 124092, 75867, 11533, 44790, + 77625, 120021, 35796, 59803, 125801, 109012, 11865, 40948, 37661, + 35162, 126232, 68770, 2305, 112100, 109343, 113652, 63116, 117632, + 34612, 30154, 52656, 21245, 80109, 63028, 109193, 13803, 81628, + 103010, 41739, 90882, 108892, 84326, 123775, 95460, 92086, 43419, + 61552, 45574, 2628, 89799, 67911, 62007, 2257, 92853, 121588, + 57810, 112838, 121329, 36503, 52739, 38030, 48557, 2670, 27008, + 56546, 89940, 94261, 56310, 42882, 105636, 34259, 26129, 17766, + 25618, 81251, 108066, 50349, 70210, 65875, 51092, 120734, 113481, + 93399, 104051, 31733, 93918, 36189, 126823, 7055, 39253, 24640, + 49923, 119661, 92746, 72370, 44090, 52579, 98340, 80882, 98169, + 113077, 55781, 6278, 49716, 86849, 29219, 113187, 82276, 120972, + 37769, 34182, 38896, 60703, 23072, 20332, 53289, 67534, 128527, + 28989, 46785, 39645, 58153, 16172, 50958, 107707, 82031, 63615, + 81865, 44843, 30274, 101009, 55711, 13111, 13384, 128841, 20243, + 48532, 81229, 79282, 19653, 117403, 86011, 23019, 91696, 68334, + 51259, 45385, 123082, 77607}; + + uint64_t R1[256] = { + 61774, 129450, 45614, 54183, 40191, 63050, 97813, 59328, 95323, + 106309, 60919, 123528, 120845, 22348, 46727, 15718, 67129, 125347, + 129458, 97385, 129338, 61503, 129444, 65283, 70304, 73753, 108456, + 40271, 123958, 110061, 101428, 97272, 115566, 46334, 56167, 38814, + 107270, 49806, 74871, 96499, 56044, 113420, 17964, 102808, 47455, + 110760, 18647, 40940, 69624, 18126, 56393, 46513, 69046, 83928, + 129932, 93391, 95276, 14422, 123309, 109213, 43303, 20241, 28811, + 50246, 50358, 121756, 74307, 83623, 72898, 41178, 38985, 27047, + 57930, 11520, 66954, 117544, 85447, 26277, 111202, 90262, 75870, + 93045, 93012, 59127, 95833, 34164, 20069, 113332, 20713, 52002, + 106065, 112067, 85093, 33622, 4518, 81688, 119200, 123522, 95456, + 17404, 68062, 117903, 127004, 40455, 1452, 129981, 9221, 28442, + 24309, 10147, 119119, 20461, 35160, 70240, 29392, 72125, 100339, + 83878, 9175, 3186, 40335, 17350, 14816, 109339, 125744, 117038, + 92864, 11912, 35015, 65176, 64650, 51875, 18980, 20466, 38058, + 110008, 101411, 74054, 34525, 71889, 2499, 99343, 118316, 35305, + 65400, 74786, 109172, 31324, 129974, 18171, 11057, 95753, 16178, + 105590, 100073, 45372, 125160, 44836, 57156, 111032, 98203, 79774, + 15089, 26366, 120665, 28298, 108590, 16104, 22881, 39249, 18974, + 28133, 16980, 126556, 12723, 48898, 35581, 27787, 82074, 123080, + 84672, 115064, 87793, 84073, 97088, 117108, 14178, 27198, 98148, + 115176, 122731, 68855, 32044, 90118, 101341, 65396, 583, 62058, + 91302, 32676, 65093, 124820, 18817, 24132, 93225, 37352, 48341, + 15503, 118581, 110828, 2844, 39350, 104648, 20734, 32894, 47882, + 9821, 94888, 104778, 123291, 120703, 122014, 40871, 55582, 87690, + 70582, 48333, 106740, 54827, 97757, 70499, 96707, 123530, 88113, + 43019, 22034, 22639, 60716, 94648, 54544, 4207}; + uint64_t One1[256] = {1}; + uint64_t r_0[256] = { + 67620, 8547, 118827, 64763, 118527, 113363, 66673, 109898, 106837, + 129342, 64436, 72034, 107436, 78839, 1683, 87839, 26180, 63620, + 82009, 21952, 67437, 72385, 10944, 35535, 124854, 101961, 92703, + 61750, 24987, 42487, 54459, 45878, 4123, 45221, 61163, 47586, + 63363, 6033, 13785, 1744, 94587, 46568, 1183, 109540, 66823, + 75163, 101942, 71193, 4460, 110617, 106285, 85350, 80563, 80676, + 130063, 2768, 124766, 121080, 54388, 56200, 82380, 108714, 94959, + 32411, 49683, 48444, 87933, 64989, 82306, 10548, 87445, 93198, + 46021, 130366, 2861, 23012, 68791, 29071, 75022, 12236, 50160, + 68371, 104611, 33286, 59566, 88040, 16599, 128457, 110725, 117076, + 31525, 117920, 41823, 104286, 71476, 85490, 8298, 63812, 18358, + 46727, 21502, 26276, 40268, 80771, 63617, 37651, 99372, 40767, + 12505, 2033, 33952, 102876, 114759, 59379, 65080, 4463, 106994, + 130597, 79101, 40214, 60487, 40128, 49183, 64405, 65778, 50446, + 98395, 1219, 4265, 56809, 32076, 83897, 55740, 51883, 99685, + 9877, 1376, 92141, 1979, 94779, 18706, 81433, 116200, 13683, + 118, 82081, 64285, 129065, 64828, 58346, 21115, 67581, 64284, + 48465, 26611, 121828, 124257, 107860, 6154, 113924, 50788, 21835, + 100579, 74327, 32774, 75268, 68575, 72216, 11777, 128300, 68955, + 72248, 98957, 52976, 36575, 92040, 102665, 26824, 76134, 116040, + 9569, 39235, 54428, 122038, 24996, 35324, 50442, 97178, 1742, + 126542, 78145, 38573, 92927, 116296, 48934, 88686, 101045, 84286, + 122142, 30108, 42954, 82007, 118419, 115100, 128624, 91060, 81389, + 4745, 59970, 98841, 19370, 7182, 120000, 52011, 17397, 78187, + 9846, 21795, 130420, 95245, 26526, 65913, 79790, 98194, 10150, + 103564, 77972, 57176, 24113, 32232, 66632, 129968, 99153, 30543, + 87250, 3975, 93393, 76193, 8643, 56190, 17450}; + + // CT(negModn) d_negModn d_twiddle twiddle + // int ArrLength=128; + + int ArrLength = 256; + int BitOfArrlength = log2(ArrLength); + int NumOfThreads = ArrLength >> 1; + size_t bytes = ArrLength * sizeof(uint64_t); + + // uint64_t dim = ArrLength; + dim3 blockDim(NumOfThreads); + dim3 gridDim(1); + + uint64_t *d_r_0; + cudaMalloc((void **)&d_r_0, bytes); + cudaMemcpy((void *)d_r_0, (void *)r_0, bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle; + cudaMalloc((void **)&d_con_twiddle, bytes); + cudaMemcpy((void *)d_con_twiddle, (void *)con_twiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_shoup; + cudaMalloc((void **)&d_con_twiddle_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_shoup, (void *)con_twiddle_shoup, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT; + cudaMalloc((void **)&d_con_twiddle_NCT, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT, (void *)con_twiddle_NCT, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_twiddle_NCT_shoup; + cudaMalloc((void **)&d_con_twiddle_NCT_shoup, bytes); + cudaMemcpy((void *)d_con_twiddle_NCT_shoup, (void *)con_twiddle_NCT_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle; + cudaMalloc((void **)&d_con_InvTwiddle, bytes); + cudaMemcpy((void *)d_con_InvTwiddle, (void *)con_InvTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_InvTwiddle_shoup; + cudaMalloc((void **)&d_con_InvTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_InvTwiddle_shoup, (void *)con_InvTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle; + cudaMalloc((void **)&d_con_ICTTwiddle, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle, (void *)con_ICTTwiddle, bytes, + cudaMemcpyHostToDevice); + + uint64_t *d_con_ICTTwiddle_shoup; + cudaMalloc((void **)&d_con_ICTTwiddle_shoup, bytes); + cudaMemcpy((void *)d_con_ICTTwiddle_shoup, (void *)con_ICTTwiddle_shoup, + bytes, cudaMemcpyHostToDevice); + + uint64_t *d_Modn; + cudaMalloc((void **)&d_Modn, bytes); + cudaMemcpy((void *)d_Modn, (void *)Modn, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_Modn, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_negModn; + cudaMalloc((void **)&d_negModn, bytes); + cudaMemcpy((void *)d_negModn, (void *)negModn, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_negModn, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + // Rone + uint64_t *d_nctR; + cudaMalloc((void **)&d_nctR, bytes); + cudaMemcpy((void *)d_nctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_nctR, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + + uint64_t *d_ctR; + cudaMalloc((void **)&d_ctR, bytes); + cudaMemcpy((void *)d_ctR, (void *)R1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_ctR, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + // IR + uint64_t *d_IRct; + cudaMalloc((void **)&d_IRct, bytes); + cudaMemcpy((void *)d_IRct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testctsample<<>>(d_IRct, ArrLength, MOD, d_con_twiddle, + d_con_twiddle_shoup); + + uint64_t *d_IRnct; + cudaMalloc((void **)&d_IRnct, bytes); + cudaMemcpy((void *)d_IRnct, (void *)One1, bytes, cudaMemcpyHostToDevice); + Testnctsample<<>>( + d_IRnct, ArrLength, MOD, d_con_twiddle_NCT, d_con_twiddle_NCT_shoup); + /* + cout<<"ct1:"<>>(d_arrX,d_arrY,d_negModn, + // d_con_NegModn_shoup,d_Modn,d_con_Modn_shoup,256,MOD + // ,d_con_twiddle,d_con_twiddle_shoup,d_con_ICTTwiddle,d_con_ICTTwiddle_shoup,d_con_twiddle_NCT,d_con_twiddle_NCT_shoup,d_con_InvTwiddle,d_con_InvTwiddle_shoup,inv,inv_shoup,d_sample); + + // ============================================================ + // 1. 分配锁页主机内存(pinned) + // ============================================================ + const size_t batch_bytes = (size_t)NUM_TESTS * ARR_LEN * sizeof(uint64_t); + uint64_t *h_c1_tilde = NULL; + uint64_t *h_m_batch = NULL; + uint64_t *h_result = NULL; + + CUDA_CHECK(cudaMallocHost(&h_c1_tilde, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_m_batch, batch_bytes)); + CUDA_CHECK(cudaMallocHost(&h_result, batch_bytes)); + + // ============================================================ + // 2. 随机填充测试数据 + // c1_tilde:[0, n²) 范围内的随机蒙哥马利密文(limb 表示) + // m_batch :小整数明文,仅低几个 limb 非零 + // ============================================================ + srand((unsigned int)time(NULL)); + + // n² 最高有效 limb 的索引和值(与 save22.cu 中的 n2_top_idx 一致) + const int n2_top_idx = 240; + const uint64_t n2_top_val = 27958ULL; + const uint64_t mask17 = (1ULL << BASE_BITS) - 1; + + for (int p = 0; p < NUM_TESTS; p++) { + uint64_t *c = h_c1_tilde + (size_t)p * ARR_LEN; + for (int j = 0; j < ARR_LEN; j++) { + if (j < n2_top_idx) + c[j] = ((uint64_t)rand() ^ ((uint64_t)rand() << 15)) & mask17; + else if (j == n2_top_idx) + c[j] = (uint64_t)rand() % n2_top_val; + else + c[j] = 0ULL; + } + } + + for (int p = 0; p < NUM_TESTS; p++) { + uint64_t *m = h_m_batch + (size_t)p * ARR_LEN; + // 明文只填低 4 个 limb,其余为 0(保证 m < n) + for (int j = 0; j < ARR_LEN; j++) m[j] = 0; + m[0] = (uint64_t)rand() & mask17; + m[1] = (uint64_t)rand() & mask17; + m[2] = (uint64_t)rand() & mask17; + m[3] = (uint64_t)rand() % 100; + } + + // ============================================================ + // 3. 分配设备内存(H2D 目标 / D2H 源) + // d_c1:密文输入,Step3 结果也写回此处 + // d_m :明文批 + // ============================================================ + uint64_t *d_c1_dev = NULL; + uint64_t *d_m_dev = NULL; + + CUDA_CHECK(cudaMalloc(&d_c1_dev, batch_bytes)); + CUDA_CHECK(cudaMalloc(&d_m_dev, batch_bytes)); + + // ============================================================ + // 4. 创建计时事件(复用于 warmup 和正式计时) + // ============================================================ + cudaEvent_t ev0, ev1; + CUDA_CHECK(cudaEventCreate(&ev0)); + CUDA_CHECK(cudaEventCreate(&ev1)); + + // ── 辅助 lambda:统一调用 paillier_addplain_subplain_batch ── + // 使用宏封装重复参数列表,避免书写错误 +#define CALL_BATCH(dc1, dm, sign) \ + paillier_addplain_subplain_batch( \ + (dc1), (dm), d_N, d_N2, d_ctR, d_nctR, d_negModn, d_con_NegModn_shoup, \ + d_Modn, d_con_Modn_shoup, MOD, d_con_twiddle, d_con_twiddle_shoup, \ + d_con_ICTTwiddle, d_con_ICTTwiddle_shoup, d_con_twiddle_NCT, \ + d_con_twiddle_NCT_shoup, d_con_InvTwiddle, d_con_InvTwiddle_shoup, inv, \ + inv_shoup, d_sample, (sign)) + + // ============================================================ + // 4a. Warm-up(不计时) + // 目的:消除首次 kernel 启动的 JIT / cubin 加载开销 + // 以及 PCIe DMA 路径的"冷启动"延迟 + // 流程:untimed H2D → compute → untimed D2H → sync + // 完成后重新上传,保证正式测试使用的是原始密文 + // ============================================================ + printf("[Warm-up] 上传数据...\n"); + CUDA_CHECK( + cudaMemcpy(d_c1_dev, h_c1_tilde, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_m_dev, h_m_batch, batch_bytes, cudaMemcpyHostToDevice)); + + printf("[Warm-up] AddPlain...\n"); + { + StepTime dummy = CALL_BATCH(d_c1_dev, d_m_dev, +1); + (void)dummy; + // untimed D2H(将 warmup 结果冲刷出 GPU,确保流水线真正跑完) + CUDA_CHECK( + cudaMemcpy(h_result, d_c1_dev, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaDeviceSynchronize()); + } + printf("[Warm-up] 完成\n\n"); + + // ============================================================ + // 5. [外部] 正式计时:AddPlain + // 步骤 A: 计时 H2D(d_c1_dev 恢复为原始 c1_tilde) + // 步骤 B: 计时 compute + // 步骤 C: 计时 D2H + // ============================================================ + // 步骤 A:计时 H2D + printf("[正式运行] AddPlain — H2D 上传...\n"); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(d_c1_dev, h_c1_tilde, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_m_dev, h_m_batch, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float h2d_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&h2d_ms, ev0, ev1)); + printf("[正式运行] H2D 完成:%.4f ms (%.2f GB/s)\n", h2d_ms, + 2.0 * batch_bytes / (h2d_ms * 1e-3) / 1e9); + + // 步骤 B:计时 compute(Step1 + Step2 + Step3) + printf("[正式运行] AddPlain GPU 计算...\n"); + StepTime t_add = CALL_BATCH(d_c1_dev, d_m_dev, +1); + + // 步骤 C:计时 D2H + printf("[正式运行] D2H 下载...\n"); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_result, d_c1_dev, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float d2h_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&d2h_ms, ev0, ev1)); + printf("[正式运行] D2H 完成:%.4f ms (%.2f GB/s)\n", d2h_ms, + (double)batch_bytes / (d2h_ms * 1e-3) / 1e9); + + print_steptime_ext("AddPlain batch(H2D/D2H 外部)", t_add, h2d_ms, d2h_ms); + + // 写验证文件(m / c1_tilde_original / result,供 Python 验证) + // write_results_to_txt("addplain_results.txt", + // h_m_batch, h_c1_tilde, h_result, N2_arr, N_arr); + + // ============================================================ + // 6. [外部] 正式计时:SubPlain + // 重新上传 h_c1_tilde(d_c1_dev 被 AddPlain 的 Step3 覆盖) + // ============================================================ + printf("\n[正式运行] SubPlain — H2D 上传...\n"); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(d_c1_dev, h_c1_tilde, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK( + cudaMemcpy(d_m_dev, h_m_batch, batch_bytes, cudaMemcpyHostToDevice)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float h2d_sub_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&h2d_sub_ms, ev0, ev1)); + printf("[正式运行] H2D 完成:%.4f ms\n", h2d_sub_ms); + + printf("[正式运行] SubPlain GPU 计算...\n"); + StepTime t_sub = CALL_BATCH(d_c1_dev, d_m_dev, -1); + + printf("[正式运行] D2H 下载...\n"); + CUDA_CHECK(cudaEventRecord(ev0)); + CUDA_CHECK( + cudaMemcpy(h_result, d_c1_dev, batch_bytes, cudaMemcpyDeviceToHost)); + CUDA_CHECK(cudaEventRecord(ev1)); + CUDA_CHECK(cudaEventSynchronize(ev1)); + float d2h_sub_ms = 0.f; + CUDA_CHECK(cudaEventElapsedTime(&d2h_sub_ms, ev0, ev1)); + printf("[正式运行] D2H 完成:%.4f ms\n", d2h_sub_ms); + + print_steptime_ext("SubPlain batch(H2D/D2H 外部)", t_sub, h2d_sub_ms, + d2h_sub_ms); + + // write_results_to_txt("subplain_results.txt", + // h_m_batch, h_c1_tilde, h_result, N2_arr, N_arr); + +#undef CALL_BATCH + + ////////////////////////////////////// + // ============================================================ + // 8. 释放资源 + // ============================================================ + cudaEventDestroy(ev0); + cudaEventDestroy(ev1); + + cudaFreeHost(h_c1_tilde); + cudaFreeHost(h_m_batch); + cudaFreeHost(h_result); + + cudaFree(d_c1_dev); + cudaFree(d_m_dev); + + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_nctR); + cudaFree(d_ctR); + cudaFree(d_N); + cudaFree(d_N2); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + cudaFree(d_sample); + + printf("\n[DONE] 测试完成。\n"); + return 0; + /* + cudaFree(d_sample); + cudaFree(d_con_twiddle); + cudaFree(d_con_twiddle_shoup); + cudaFree(d_con_twiddle_NCT); + cudaFree(d_con_twiddle_NCT_shoup); + cudaFree(d_con_InvTwiddle); + cudaFree(d_con_InvTwiddle_shoup); + cudaFree(d_con_ICTTwiddle); + cudaFree(d_con_ICTTwiddle_shoup); + cudaFree(d_Modn); + cudaFree(d_negModn); + cudaFree(d_con_Modn_shoup); + cudaFree(d_con_NegModn_shoup); + + + //free(arrX); + //free(arrY); +*/ + // 重置设备状态 (可选,但推荐,方便性能分析工具 profiling) + cudaDeviceReset(); +}